From 4c87c9f14f6e3c7f60e80ce6ec4575a3fbd9e9db Mon Sep 17 00:00:00 2001 From: Shelikhoo Date: Thu, 19 Feb 2026 20:01:29 +0000 Subject: [PATCH] Add Wireguard Support for v2ray (mostly vibe coded) --- common/packetswitch/gvisorstack/adapter.go | 246 ++++++++++ .../packetswitch/gvisorstack/adapter_test.go | 458 ++++++++++++++++++ common/packetswitch/gvisorstack/config.pb.go | 190 ++++++++ common/packetswitch/gvisorstack/config.proto | 22 + common/packetswitch/gvisorstack/dialer.go | 130 +++++ common/packetswitch/gvisorstack/stack.go | 164 +++++++ .../packetswitch/interconnect/interconnect.go | 3 + .../interconnect/networkLayer_cable.go | 92 ++++ .../interconnect/networkLayer_cable_test.go | 182 +++++++ common/packetswitch/packetswitch.go | 17 + go.mod | 2 + go.sum | 4 + main/distro/all/all.go | 3 + proxy/wireguard/outbound/config.pb.go | 216 +++++++++ proxy/wireguard/outbound/config.proto | 31 ++ proxy/wireguard/outbound/errors.generated.go | 9 + proxy/wireguard/outbound/outbound.go | 387 +++++++++++++++ proxy/wireguard/wgcommon/config.pb.go | 233 +++++++++ proxy/wireguard/wgcommon/config.proto | 24 + proxy/wireguard/wgcommon/errors.generated.go | 9 + proxy/wireguard/wgcommon/setup.go | 121 +++++ proxy/wireguard/wgcommon/wgConnAdaptor.go | 181 +++++++ .../wireguard/wgcommon/wgConnAdaptor_test.go | 250 ++++++++++ proxy/wireguard/wgcommon/wgDeviceAdaptor.go | 201 ++++++++ .../wgcommon/wgDeviceAdaptor_test.go | 276 +++++++++++ proxy/wireguard/wgcommon/wgLogAdaptor.go | 27 ++ proxy/wireguard/wgcommon/wgcommon.go | 3 + proxy/wireguard/wgcommon/wgdevice.go | 71 +++ .../clicommand/enrollmentlink_cli.go | 1 + .../mirrorenrollment/enrollmentlink.go | 2 + 30 files changed, 3555 insertions(+) create mode 100644 common/packetswitch/gvisorstack/adapter.go create mode 100644 common/packetswitch/gvisorstack/adapter_test.go create mode 100644 common/packetswitch/gvisorstack/config.pb.go create mode 100644 common/packetswitch/gvisorstack/config.proto create mode 100644 common/packetswitch/gvisorstack/dialer.go create mode 100644 common/packetswitch/gvisorstack/stack.go create mode 100644 common/packetswitch/interconnect/interconnect.go create mode 100644 common/packetswitch/interconnect/networkLayer_cable.go create mode 100644 common/packetswitch/interconnect/networkLayer_cable_test.go create mode 100644 common/packetswitch/packetswitch.go create mode 100644 proxy/wireguard/outbound/config.pb.go create mode 100644 proxy/wireguard/outbound/config.proto create mode 100644 proxy/wireguard/outbound/errors.generated.go create mode 100644 proxy/wireguard/outbound/outbound.go create mode 100644 proxy/wireguard/wgcommon/config.pb.go create mode 100644 proxy/wireguard/wgcommon/config.proto create mode 100644 proxy/wireguard/wgcommon/errors.generated.go create mode 100644 proxy/wireguard/wgcommon/setup.go create mode 100644 proxy/wireguard/wgcommon/wgConnAdaptor.go create mode 100644 proxy/wireguard/wgcommon/wgConnAdaptor_test.go create mode 100644 proxy/wireguard/wgcommon/wgDeviceAdaptor.go create mode 100644 proxy/wireguard/wgcommon/wgDeviceAdaptor_test.go create mode 100644 proxy/wireguard/wgcommon/wgLogAdaptor.go create mode 100644 proxy/wireguard/wgcommon/wgcommon.go create mode 100644 proxy/wireguard/wgcommon/wgdevice.go diff --git a/common/packetswitch/gvisorstack/adapter.go b/common/packetswitch/gvisorstack/adapter.go new file mode 100644 index 000000000..76a19e744 --- /dev/null +++ b/common/packetswitch/gvisorstack/adapter.go @@ -0,0 +1,246 @@ +package gvisorstack + +import ( + "context" + "sync" + + "github.com/v2fly/v2ray-core/v5/common" + "github.com/v2fly/v2ray-core/v5/common/packetswitch" + "gvisor.dev/gvisor/pkg/buffer" + "gvisor.dev/gvisor/pkg/tcpip" + "gvisor.dev/gvisor/pkg/tcpip/header" + "gvisor.dev/gvisor/pkg/tcpip/network/ipv4" + "gvisor.dev/gvisor/pkg/tcpip/network/ipv6" + "gvisor.dev/gvisor/pkg/tcpip/stack" +) + +func NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(ctx context.Context, mtu int, networkLayerSwitch packetswitch.NetworkLayerDevice) *NetworkLayerDeviceToGvisorLinkEndpointAdaptor { + return &NetworkLayerDeviceToGvisorLinkEndpointAdaptor{ + mtu: mtu, + networkLayerSwitch: networkLayerSwitch, + waitCh: make(chan struct{}), + } +} + +// NetworkLayerDeviceToGvisorLinkEndpointAdaptor is primarily machine generated. +type NetworkLayerDeviceToGvisorLinkEndpointAdaptor struct { + mtu int + networkLayerSwitch packetswitch.NetworkLayerDevice + + mu sync.RWMutex + dispatcher stack.NetworkDispatcher + attached bool + closed bool + onClose func() + waitCh chan struct{} +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) MTU() uint32 { + n.mu.RLock() + defer n.mu.RUnlock() + return uint32(n.mtu) +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) SetMTU(mtu uint32) { + n.mu.Lock() + defer n.mu.Unlock() + n.mtu = int(mtu) +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) MaxHeaderLength() uint16 { + // No additional link-layer header. + return 0 +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) LinkAddress() tcpip.LinkAddress { + // Not applicable for network-layer device. + return "" +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) SetLinkAddress(addr tcpip.LinkAddress) { + // no-op +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) Capabilities() stack.LinkEndpointCapabilities { + return stack.CapabilityNone +} + +// networkLayerWriter adapts packets from NetworkLayerDevice into gVisor Stack. +type networkLayerWriter struct { + parent *NetworkLayerDeviceToGvisorLinkEndpointAdaptor +} + +func (w *networkLayerWriter) Write(packet []byte) (int, error) { + if len(packet) == 0 { + return 0, nil + } + + buf := buffer.MakeWithData(packet) + pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ + Payload: buf, + // Do not call buf.Release here; PacketBuffer.DecRef will release internal buffer. + }) + + // Determine network protocol by IP version. + ver := packet[0] >> 4 + var proto tcpip.NetworkProtocolNumber + switch ver { + case 4: + proto = ipv4.ProtocolNumber + case 6: + proto = ipv6.ProtocolNumber + default: + // Unknown network packet, drop. + pkt.DecRef() + return 0, nil + } + + w.parent.mu.RLock() + d := w.parent.dispatcher + w.parent.mu.RUnlock() + if d == nil { + // No dispatcher attached, drop. + pkt.DecRef() + return 0, nil + } + + // Deliver to network layer. The dispatcher takes ownership of pkt + // and is responsible for releasing it. + d.DeliverNetworkPacket(proto, pkt) + return len(packet), nil +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) Attach(dispatcher stack.NetworkDispatcher) { + n.mu.Lock() + defer n.mu.Unlock() + if dispatcher == nil { + // Detaching. + n.dispatcher = nil + n.attached = false + return + } + + n.dispatcher = dispatcher + writer := &networkLayerWriter{parent: n} + // Let the network layer device know where to write incoming packets. + if err := n.networkLayerSwitch.OnAttach(writer); err == nil { + n.attached = true + } else { + // OnAttach failed; keep attached false. + n.attached = false + } +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) IsAttached() bool { + n.mu.RLock() + defer n.mu.RUnlock() + return n.attached +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) Wait() { + // If closed, return immediately. + n.mu.RLock() + closed := n.closed + ch := n.waitCh + n.mu.RUnlock() + if closed { + return + } + // Wait until closed is signaled. + <-ch +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) ARPHardwareType() header.ARPHardwareType { + return header.ARPHardwareNone +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) AddHeader(buffer *stack.PacketBuffer) { + // No link-layer header to add. +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) ParseHeader(buffer *stack.PacketBuffer) bool { + // Nothing to parse; packet is a bare network packet. + return true +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) Close() { + n.mu.Lock() + if n.closed { + n.mu.Unlock() + return + } + n.closed = true + n.attached = false + n.mu.Unlock() + + // Close underlying network device if any. + _ = common.Close(n.networkLayerSwitch) + + // Run onClose action if set. + n.mu.RLock() + onc := n.onClose + ch := n.waitCh + n.mu.RUnlock() + if onc != nil { + onc() + } + + // Signal waiters. + close(ch) +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) SetOnCloseAction(f func()) { + n.mu.Lock() + defer n.mu.Unlock() + n.onClose = f +} + +func (n *NetworkLayerDeviceToGvisorLinkEndpointAdaptor) WritePackets(list stack.PacketBufferList) (int, tcpip.Error) { + // Defensive: if receiver is nil, treat as closed. + if n == nil { + return 0, &tcpip.ErrClosedForSend{} + } + // Convert each packet to bytes and write to networkLayerSwitch. + slice := list.AsSlice() + if len(slice) == 0 { + return 0, nil + } + + n.mu.RLock() + dev := n.networkLayerSwitch + mtu := n.mtu + n.mu.RUnlock() + if dev == nil { + return 0, &tcpip.ErrClosedForSend{} + } + + written := 0 + for _, pkt := range slice { + if pkt == nil { + continue + } + // Get slices and copy into a contiguous buffer. + slices := pkt.AsSlices() + total := 0 + for _, s := range slices { + total += len(s) + } + if mtu > 0 && total > mtu { + return written, &tcpip.ErrMessageTooLong{} + } + cp := make([]byte, total) + off := 0 + for _, s := range slices { + copy(cp[off:], s) + off += len(s) + } + _, err := dev.Write(cp) + if err != nil { + // Map writer error to tcpip error. + return written, &tcpip.ErrNoBufferSpace{} + } + written++ + } + + return written, nil +} diff --git a/common/packetswitch/gvisorstack/adapter_test.go b/common/packetswitch/gvisorstack/adapter_test.go new file mode 100644 index 000000000..741fc6f0e --- /dev/null +++ b/common/packetswitch/gvisorstack/adapter_test.go @@ -0,0 +1,458 @@ +package gvisorstack + +import ( + "context" + "errors" + "reflect" + "sync" + "testing" + "time" + + "github.com/v2fly/v2ray-core/v5/common/packetswitch" + "gvisor.dev/gvisor/pkg/buffer" + "gvisor.dev/gvisor/pkg/tcpip" + "gvisor.dev/gvisor/pkg/tcpip/network/ipv4" + "gvisor.dev/gvisor/pkg/tcpip/network/ipv6" + "gvisor.dev/gvisor/pkg/tcpip/stack" +) + +// fakeDevice implements packetswitch.NetworkLayerDevice for testing. +type fakeDevice struct { + mu sync.Mutex + writer packetswitch.NetworkLayerPacketWriter + writes [][]byte + closed bool + onAttach func(packetswitch.NetworkLayerPacketWriter) error +} + +func (f *fakeDevice) OnAttach(w packetswitch.NetworkLayerPacketWriter) error { + f.mu.Lock() + defer f.mu.Unlock() + f.writer = w + if f.onAttach != nil { + return f.onAttach(w) + } + return nil +} + +func (f *fakeDevice) Write(packet []byte) (int, error) { + f.mu.Lock() + defer f.mu.Unlock() + if f.closed { + return 0, errors.New("closed") + } + // make a copy to avoid aliasing test buffer. + cp := make([]byte, len(packet)) + copy(cp, packet) + f.writes = append(f.writes, cp) + return len(packet), nil +} + +func (f *fakeDevice) Close() error { + f.mu.Lock() + defer f.mu.Unlock() + f.closed = true + return nil +} + +func (f *fakeDevice) getWriter() packetswitch.NetworkLayerPacketWriter { + f.mu.Lock() + defer f.mu.Unlock() + return f.writer +} + +func (f *fakeDevice) lastWrite() []byte { + f.mu.Lock() + defer f.mu.Unlock() + if len(f.writes) == 0 { + return nil + } + return f.writes[len(f.writes)-1] +} + +// fakeDispatcher implements stack.NetworkDispatcher, capturing delivered packets. +type fakeDispatcher struct { + mu sync.Mutex + protocols []tcpip.NetworkProtocolNumber + pkts [][]byte +} + +func (d *fakeDispatcher) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { + // capture packet payload safely by copying from AsSlices + slices := pkt.AsSlices() + total := 0 + for _, s := range slices { + total += len(s) + } + cp := make([]byte, total) + off := 0 + for _, s := range slices { + copy(cp[off:], s) + off += len(s) + } + // record + d.mu.Lock() + d.protocols = append(d.protocols, protocol) + d.pkts = append(d.pkts, cp) + d.mu.Unlock() + // release the packet + pkt.DecRef() +} + +func (d *fakeDispatcher) DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { + // not used in these tests + pkt.DecRef() +} + +func (d *fakeDispatcher) last() (tcpip.NetworkProtocolNumber, []byte) { + d.mu.Lock() + defer d.mu.Unlock() + if len(d.protocols) == 0 { + return 0, nil + } + return d.protocols[len(d.protocols)-1], d.pkts[len(d.pkts)-1] +} + +func TestAttachAndInboundIPv4IPv6(t *testing.T) { + dev := &fakeDevice{} + a := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(context.Background(), 1500, dev) + d := &fakeDispatcher{} + // Initially not attached + if a.IsAttached() { + t.Fatal("expected not attached") + } + + // Attach should call device.OnAttach and store writer + a.Attach(d) + w := dev.getWriter() + if w == nil { + t.Fatal("device did not receive writer on attach") + } + if !a.IsAttached() { + t.Fatal("expected attached after successful OnAttach") + } + + // Send an IPv4 packet (first byte 0x45 = version 4, IHL 5) + ipv4pkt := []byte{0x45, 0x00, 0x00, 0x04} + if _, err := w.Write(ipv4pkt); err != nil { + t.Fatalf("writer.Write failed: %v", err) + } + // Check dispatcher received it + proto, payload := d.last() + if proto != ipv4.ProtocolNumber { + t.Fatalf("expected ipv4 protocol, got %v", proto) + } + if !reflect.DeepEqual(payload, ipv4pkt) { + t.Fatalf("unexpected payload: %v", payload) + } + + // IPv6 packet (first byte 0x60) + ipv6pkt := []byte{0x60, 0x00, 0x00, 0x00} + if _, err := w.Write(ipv6pkt); err != nil { + t.Fatalf("writer.Write failed: %v", err) + } + proto2, payload2 := d.last() + if proto2 != ipv6.ProtocolNumber { + t.Fatalf("expected ipv6 protocol, got %v", proto2) + } + if !reflect.DeepEqual(payload2, ipv6pkt) { + t.Fatalf("unexpected payload: %v", payload2) + } +} + +func TestInboundNonIPIsDropped(t *testing.T) { + dev := &fakeDevice{} + a := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(context.Background(), 1500, dev) + d := &fakeDispatcher{} + a.Attach(d) + w := dev.getWriter() + if w == nil { + t.Fatal("device did not receive writer on attach") + } + + // Non-IP packet: first nibble 0 + pkt := []byte{0x00, 0x01, 0x02} + if _, err := w.Write(pkt); err != nil { + t.Fatalf("writer.Write failed: %v", err) + } + // Dispatcher should not have any packets + proto, payload := d.last() + if payload != nil || proto != 0 { + t.Fatalf("expected no delivery for non-ip packet, got proto=%v payload=%v", proto, payload) + } +} + +func makePacketBufferPayload(b []byte) *stack.PacketBuffer { + buf := buffer.MakeWithData(b) + pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ + Payload: buf, + }) + return pkt +} + +func TestOutboundWritePacketsOk(t *testing.T) { + dev := &fakeDevice{} + a := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(context.Background(), 1500, dev) + + // Prepare a PacketBufferList with one packet + list := stack.PacketBufferList{} + payload := []byte{0x45, 0x01, 0x02, 0x03} + pkt := makePacketBufferPayload(payload) + list.PushBack(pkt) + + written, err := a.WritePackets(list) + if err != nil { + t.Fatalf("WritePackets returned error: %v", err) + } + if written != 1 { + t.Fatalf("expected 1 written, got %d", written) + } + // Device should have received the payload + lw := dev.lastWrite() + if !reflect.DeepEqual(lw, payload) { + t.Fatalf("device write mismatch: got %v, want %v", lw, payload) + } +} + +func TestWritePacketsMTUExceeded(t *testing.T) { + dev := &fakeDevice{} + // mtu set to 2 + a := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(context.Background(), 2, dev) + + list := stack.PacketBufferList{} + payload := []byte{0x45, 0x01, 0x02} // len 3 > mtu 2 + pkt := makePacketBufferPayload(payload) + list.PushBack(pkt) + + written, err := a.WritePackets(list) + if err == nil { + t.Fatalf("expected error due to message too long, got nil") + } + if _, ok := err.(*tcpip.ErrMessageTooLong); !ok { + t.Fatalf("expected ErrMessageTooLong, got %T", err) + } + if written != 0 { + t.Fatalf("expected 0 written, got %d", written) + } +} + +func TestCloseAndOnCloseActionAndWait(t *testing.T) { + dev := &fakeDevice{} + a := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(context.Background(), 1500, dev) + called := make(chan struct{}) + a.SetOnCloseAction(func() { close(called) }) + + // Wait should block until Close is called. Run Wait in goroutine. + done := make(chan struct{}) + go func() { + a.Wait() + close(done) + }() + + // Give goroutine a moment to start + time.Sleep(5 * time.Millisecond) + // Now close + a.Close() + + // onClose action should be called + select { + case <-called: + // ok + case <-time.After(100 * time.Millisecond): + t.Fatal("onClose was not called") + } + + // Wait should return + select { + case <-done: + // ok + case <-time.After(100 * time.Millisecond): + t.Fatal("Wait did not return after Close") + } +} + +func TestAttachOnAttachFailLeavesNotAttached(t *testing.T) { + dev := &fakeDevice{} + // make OnAttach return error + dev.onAttach = func(w packetswitch.NetworkLayerPacketWriter) error { + return errors.New("attach fail") + } + a := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(context.Background(), 1500, dev) + d := &fakeDispatcher{} + a.Attach(d) + if a.IsAttached() { + t.Fatal("expected not attached when OnAttach fails") + } +} + +func TestWritePacketsWhenNoDevice(t *testing.T) { + // Create adaptor with nil device + a := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(context.Background(), 1500, nil) + list := stack.PacketBufferList{} + payload := []byte{0x45} + pkt := makePacketBufferPayload(payload) + list.PushBack(pkt) + written, err := a.WritePackets(list) + if err == nil { + t.Fatalf("expected ErrClosedForSend, got nil") + } + if _, ok := err.(*tcpip.ErrClosedForSend); !ok { + t.Fatalf("expected ErrClosedForSend, got %T", err) + } + if written != 0 { + t.Fatalf("expected 0 written, got %d", written) + } +} + +func TestSetMTUAndCapsAndHeaders(t *testing.T) { + // Create adaptor and test MTU setter + a := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(context.Background(), 1500, nil) + if a.MTU() != 1500 { + t.Fatalf("initial MTU mismatch: %d", a.MTU()) + } + a.SetMTU(9000) + if a.MTU() != 9000 { + t.Fatalf("MTU not updated: %d", a.MTU()) + } + + // Caps and headers + if a.MaxHeaderLength() != 0 { + t.Fatalf("MaxHeaderLength expected 0, got %d", a.MaxHeaderLength()) + } + if a.Capabilities() != stack.CapabilityNone { + t.Fatalf("Capabilities expected CapabilityNone, got %v", a.Capabilities()) + } + if a.LinkAddress() != "" { + t.Fatalf("LinkAddress expected empty, got %v", a.LinkAddress()) + } + // AddHeader/ParseHeader should not panic and parse returns true + pkt := makePacketBufferPayload([]byte{0x45}) + // Should be no-op + a.AddHeader(pkt) + if !a.ParseHeader(pkt) { + t.Fatalf("ParseHeader expected true") + } + pkt.DecRef() +} + +// errDevice fails after a certain number of writes to simulate partial failures. +type errDevice struct { + mu sync.Mutex + writer packetswitch.NetworkLayerPacketWriter + writes [][]byte + failAfter int + calls int + closed bool +} + +func (e *errDevice) OnAttach(w packetswitch.NetworkLayerPacketWriter) error { + e.mu.Lock() + defer e.mu.Unlock() + e.writer = w + return nil +} + +func (e *errDevice) Write(packet []byte) (int, error) { + e.mu.Lock() + defer e.mu.Unlock() + if e.closed { + return 0, errors.New("closed") + } + e.calls++ + if e.failAfter > 0 && e.calls > e.failAfter { + return 0, errors.New("injected write error") + } + cp := make([]byte, len(packet)) + copy(cp, packet) + e.writes = append(e.writes, cp) + return len(packet), nil +} + +func (e *errDevice) Close() error { + e.mu.Lock() + defer e.mu.Unlock() + e.closed = true + return nil +} + +func TestWritePacketsPartialOnError(t *testing.T) { + e := &errDevice{failAfter: 1} + a := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(context.Background(), 1500, e) + + list := stack.PacketBufferList{} + p1 := makePacketBufferPayload([]byte{0x45, 0x01}) + p2 := makePacketBufferPayload([]byte{0x45, 0x02}) + list.PushBack(p1) + list.PushBack(p2) + + written, err := a.WritePackets(list) + if err == nil { + t.Fatalf("expected error due to injected write error") + } + // We mapped device write errors to ErrNoBufferSpace + if _, ok := err.(*tcpip.ErrNoBufferSpace); !ok { + t.Fatalf("expected ErrNoBufferSpace, got %T", err) + } + if written != 1 { + t.Fatalf("expected 1 written, got %d", written) + } +} + +func TestMultiplePacketsWrite(t *testing.T) { + dev := &fakeDevice{} + a := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(context.Background(), 1500, dev) + + list := stack.PacketBufferList{} + p1 := makePacketBufferPayload([]byte{0x45, 0x01}) + p2 := makePacketBufferPayload([]byte{0x45, 0x02}) + p3 := makePacketBufferPayload([]byte{0x45, 0x03}) + list.PushBack(p1) + list.PushBack(p2) + list.PushBack(p3) + + written, err := a.WritePackets(list) + if err != nil { + t.Fatalf("WritePackets returned error: %v", err) + } + if written != 3 { + t.Fatalf("expected 3 written, got %d", written) + } + // last write should be p3 + lw := dev.lastWrite() + if !reflect.DeepEqual(lw, []byte{0x45, 0x03}) { + t.Fatalf("unexpected last write: %v", lw) + } +} + +func TestConcurrentWritePacketsAndClose(t *testing.T) { + dev := &fakeDevice{} + a := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(context.Background(), 1500, dev) + + wg := sync.WaitGroup{} + nWorkers := 5 + nPerWorker := 50 + wg.Add(nWorkers) + + for i := 0; i < nWorkers; i++ { + go func(id int) { + defer wg.Done() + for j := 0; j < nPerWorker; j++ { + list := stack.PacketBufferList{} + payload := []byte{0x45, byte(id), byte(j)} + pkt := makePacketBufferPayload(payload) + list.PushBack(pkt) + _, _ = a.WritePackets(list) + } + }(i) + } + + // Close after a short delay + go func() { + time.Sleep(10 * time.Millisecond) + a.Close() + }() + + wg.Wait() + // Wait should return quickly after Close + a.Wait() +} diff --git a/common/packetswitch/gvisorstack/config.pb.go b/common/packetswitch/gvisorstack/config.pb.go new file mode 100644 index 000000000..9563b8866 --- /dev/null +++ b/common/packetswitch/gvisorstack/config.pb.go @@ -0,0 +1,190 @@ +package gvisorstack + +import ( + routercommon "github.com/v2fly/v2ray-core/v5/app/router/routercommon" + _ "github.com/v2fly/v2ray-core/v5/common/protoext" + internet "github.com/v2fly/v2ray-core/v5/transport/internet" + protoreflect "google.golang.org/protobuf/reflect/protoreflect" + protoimpl "google.golang.org/protobuf/runtime/protoimpl" + reflect "reflect" + sync "sync" + unsafe "unsafe" +) + +const ( + // Verify that this generated code is sufficiently up-to-date. + _ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion) + // Verify that runtime/protoimpl is sufficiently up-to-date. + _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) +) + +type Config struct { + state protoimpl.MessageState `protogen:"open.v1"` + Mtu uint32 `protobuf:"varint,2,opt,name=mtu,proto3" json:"mtu,omitempty"` + UserLevel uint32 `protobuf:"varint,3,opt,name=user_level,json=userLevel,proto3" json:"user_level,omitempty"` + Ips []*routercommon.CIDR `protobuf:"bytes,6,rep,name=ips,proto3" json:"ips,omitempty"` + Routes []*routercommon.CIDR `protobuf:"bytes,7,rep,name=routes,proto3" json:"routes,omitempty"` + EnablePromiscuousMode bool `protobuf:"varint,8,opt,name=enable_promiscuous_mode,json=enablePromiscuousMode,proto3" json:"enable_promiscuous_mode,omitempty"` + EnableSpoofing bool `protobuf:"varint,9,opt,name=enable_spoofing,json=enableSpoofing,proto3" json:"enable_spoofing,omitempty"` + SocketSettings *internet.SocketConfig `protobuf:"bytes,10,opt,name=socket_settings,json=socketSettings,proto3" json:"socket_settings,omitempty"` + PreferIpv6ForUdp bool `protobuf:"varint,11,opt,name=prefer_ipv6_for_udp,json=preferIpv6ForUdp,proto3" json:"prefer_ipv6_for_udp,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *Config) Reset() { + *x = Config{} + mi := &file_common_packetswitch_gvisorstack_config_proto_msgTypes[0] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *Config) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*Config) ProtoMessage() {} + +func (x *Config) ProtoReflect() protoreflect.Message { + mi := &file_common_packetswitch_gvisorstack_config_proto_msgTypes[0] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use Config.ProtoReflect.Descriptor instead. +func (*Config) Descriptor() ([]byte, []int) { + return file_common_packetswitch_gvisorstack_config_proto_rawDescGZIP(), []int{0} +} + +func (x *Config) GetMtu() uint32 { + if x != nil { + return x.Mtu + } + return 0 +} + +func (x *Config) GetUserLevel() uint32 { + if x != nil { + return x.UserLevel + } + return 0 +} + +func (x *Config) GetIps() []*routercommon.CIDR { + if x != nil { + return x.Ips + } + return nil +} + +func (x *Config) GetRoutes() []*routercommon.CIDR { + if x != nil { + return x.Routes + } + return nil +} + +func (x *Config) GetEnablePromiscuousMode() bool { + if x != nil { + return x.EnablePromiscuousMode + } + return false +} + +func (x *Config) GetEnableSpoofing() bool { + if x != nil { + return x.EnableSpoofing + } + return false +} + +func (x *Config) GetSocketSettings() *internet.SocketConfig { + if x != nil { + return x.SocketSettings + } + return nil +} + +func (x *Config) GetPreferIpv6ForUdp() bool { + if x != nil { + return x.PreferIpv6ForUdp + } + return false +} + +var File_common_packetswitch_gvisorstack_config_proto protoreflect.FileDescriptor + +const file_common_packetswitch_gvisorstack_config_proto_rawDesc = "" + + "\n" + + ",common/packetswitch/gvisorstack/config.proto\x12*v2ray.core.common.packetswitch.gvisorstack\x1a$app/router/routercommon/common.proto\x1a\x1ftransport/internet/config.proto\x1a common/protoext/extensions.proto\"\x9d\x03\n" + + "\x06Config\x12\x10\n" + + "\x03mtu\x18\x02 \x01(\rR\x03mtu\x12\x1d\n" + + "\n" + + "user_level\x18\x03 \x01(\rR\tuserLevel\x12:\n" + + "\x03ips\x18\x06 \x03(\v2(.v2ray.core.app.router.routercommon.CIDRR\x03ips\x12@\n" + + "\x06routes\x18\a \x03(\v2(.v2ray.core.app.router.routercommon.CIDRR\x06routes\x126\n" + + "\x17enable_promiscuous_mode\x18\b \x01(\bR\x15enablePromiscuousMode\x12'\n" + + "\x0fenable_spoofing\x18\t \x01(\bR\x0eenableSpoofing\x12T\n" + + "\x0fsocket_settings\x18\n" + + " \x01(\v2+.v2ray.core.transport.internet.SocketConfigR\x0esocketSettings\x12-\n" + + "\x13prefer_ipv6_for_udp\x18\v \x01(\bR\x10preferIpv6ForUdpB\x9f\x01\n" + + ".com.v2ray.core.common.packetswitch.gvisorstackP\x01Z>github.com/v2fly/v2ray-core/v5/common/packetswitch/gvisorstack\xaa\x02*V2Ray.Core.Common.Packetswitch.Gvisorstackb\x06proto3" + +var ( + file_common_packetswitch_gvisorstack_config_proto_rawDescOnce sync.Once + file_common_packetswitch_gvisorstack_config_proto_rawDescData []byte +) + +func file_common_packetswitch_gvisorstack_config_proto_rawDescGZIP() []byte { + file_common_packetswitch_gvisorstack_config_proto_rawDescOnce.Do(func() { + file_common_packetswitch_gvisorstack_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_common_packetswitch_gvisorstack_config_proto_rawDesc), len(file_common_packetswitch_gvisorstack_config_proto_rawDesc))) + }) + return file_common_packetswitch_gvisorstack_config_proto_rawDescData +} + +var file_common_packetswitch_gvisorstack_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1) +var file_common_packetswitch_gvisorstack_config_proto_goTypes = []any{ + (*Config)(nil), // 0: v2ray.core.common.packetswitch.gvisorstack.Config + (*routercommon.CIDR)(nil), // 1: v2ray.core.app.router.routercommon.CIDR + (*internet.SocketConfig)(nil), // 2: v2ray.core.transport.internet.SocketConfig +} +var file_common_packetswitch_gvisorstack_config_proto_depIdxs = []int32{ + 1, // 0: v2ray.core.common.packetswitch.gvisorstack.Config.ips:type_name -> v2ray.core.app.router.routercommon.CIDR + 1, // 1: v2ray.core.common.packetswitch.gvisorstack.Config.routes:type_name -> v2ray.core.app.router.routercommon.CIDR + 2, // 2: v2ray.core.common.packetswitch.gvisorstack.Config.socket_settings:type_name -> v2ray.core.transport.internet.SocketConfig + 3, // [3:3] is the sub-list for method output_type + 3, // [3:3] is the sub-list for method input_type + 3, // [3:3] is the sub-list for extension type_name + 3, // [3:3] is the sub-list for extension extendee + 0, // [0:3] is the sub-list for field type_name +} + +func init() { file_common_packetswitch_gvisorstack_config_proto_init() } +func file_common_packetswitch_gvisorstack_config_proto_init() { + if File_common_packetswitch_gvisorstack_config_proto != nil { + return + } + type x struct{} + out := protoimpl.TypeBuilder{ + File: protoimpl.DescBuilder{ + GoPackagePath: reflect.TypeOf(x{}).PkgPath(), + RawDescriptor: unsafe.Slice(unsafe.StringData(file_common_packetswitch_gvisorstack_config_proto_rawDesc), len(file_common_packetswitch_gvisorstack_config_proto_rawDesc)), + NumEnums: 0, + NumMessages: 1, + NumExtensions: 0, + NumServices: 0, + }, + GoTypes: file_common_packetswitch_gvisorstack_config_proto_goTypes, + DependencyIndexes: file_common_packetswitch_gvisorstack_config_proto_depIdxs, + MessageInfos: file_common_packetswitch_gvisorstack_config_proto_msgTypes, + }.Build() + File_common_packetswitch_gvisorstack_config_proto = out.File + file_common_packetswitch_gvisorstack_config_proto_goTypes = nil + file_common_packetswitch_gvisorstack_config_proto_depIdxs = nil +} diff --git a/common/packetswitch/gvisorstack/config.proto b/common/packetswitch/gvisorstack/config.proto new file mode 100644 index 000000000..cd154b5d6 --- /dev/null +++ b/common/packetswitch/gvisorstack/config.proto @@ -0,0 +1,22 @@ +syntax = "proto3"; + +package v2ray.core.common.packetswitch.gvisorstack; +option csharp_namespace = "V2Ray.Core.Common.Packetswitch.Gvisorstack"; +option go_package = "github.com/v2fly/v2ray-core/v5/common/packetswitch/gvisorstack"; +option java_package = "com.v2ray.core.common.packetswitch.gvisorstack"; +option java_multiple_files = true; + +import "app/router/routercommon/common.proto"; +import "transport/internet/config.proto"; +import "common/protoext/extensions.proto"; + +message Config { + uint32 mtu = 2; + uint32 user_level = 3; + repeated v2ray.core.app.router.routercommon.CIDR ips = 6; + repeated v2ray.core.app.router.routercommon.CIDR routes = 7; + bool enable_promiscuous_mode = 8; + bool enable_spoofing = 9; + v2ray.core.transport.internet.SocketConfig socket_settings = 10; + bool prefer_ipv6_for_udp = 11; +} \ No newline at end of file diff --git a/common/packetswitch/gvisorstack/dialer.go b/common/packetswitch/gvisorstack/dialer.go new file mode 100644 index 000000000..022a39bde --- /dev/null +++ b/common/packetswitch/gvisorstack/dialer.go @@ -0,0 +1,130 @@ +package gvisorstack + +import ( + "context" + "fmt" + + "gvisor.dev/gvisor/pkg/tcpip" + "gvisor.dev/gvisor/pkg/tcpip/adapters/gonet" + "gvisor.dev/gvisor/pkg/tcpip/network/ipv4" + "gvisor.dev/gvisor/pkg/tcpip/network/ipv6" + + "github.com/v2fly/v2ray-core/v5/common/net" +) + +// DialTCP will create a connection to the given destination, using the stack. +// Machine Generated +func (w *WrappedStack) DialTCP(ctx context.Context, remoteAddress net.Destination) (net.Conn, error) { + if w == nil || w.stack == nil { + return nil, fmt.Errorf("gvisor stack not initialized") + } + + if remoteAddress.Network != net.Network_TCP { + return nil, fmt.Errorf("destination is not tcp: %v", remoteAddress.Network) + } + + // Resolve address to IP if necessary. + var ipBytes []byte + switch remoteAddress.Address.Family() { + case net.AddressFamilyIPv4: + ipBytes = remoteAddress.Address.IP().To4() + case net.AddressFamilyIPv6: + ipBytes = remoteAddress.Address.IP().To16() + case net.AddressFamilyDomain: + // Do not resolve domain names here. Return explicit error. + return nil, fmt.Errorf("domain address not supported for gVisor dial: %s", remoteAddress.Address.String()) + default: + return nil, fmt.Errorf("unsupported address family: %v", remoteAddress.Address.Family()) + } + + if ipBytes == nil { + return nil, fmt.Errorf("failed to obtain IP bytes for %v", remoteAddress.Address) + } + + // Choose network protocol number based on IP length. + netProto := ipv4.ProtocolNumber + if len(ipBytes) == 16 { + netProto = ipv6.ProtocolNumber + } + + remote := tcpip.FullAddress{ + Addr: tcpip.AddrFromSlice(ipBytes), + Port: uint16(remoteAddress.Port), + } + + // Use gonet dialer to create a TCP connection on the in-memory stack. + conn, err := gonet.DialContextTCP(ctx, w.stack, remote, netProto) + if err != nil { + return nil, err + } + return conn, nil +} + +// ListenUDP will create a connection to the given destination, using the stack. +// Machine Generated +func (w *WrappedStack) ListenUDP(ctx context.Context, localAddress net.Destination) (net.PacketConn, error) { + // allow ctx to be accepted by the function signature + _ = ctx + + if w == nil || w.stack == nil { + return nil, fmt.Errorf("gvisor stack not initialized") + } + + if localAddress.Network != net.Network_UDP { + return nil, fmt.Errorf("destination is not udp: %v", localAddress.Network) + } + + // Determine local address bytes. + var ipBytes []byte + specified := false + if localAddress.Address == nil { + // If address is nil, treat as unspecified (zero) address. + specified = false + } else { + switch localAddress.Address.Family() { + case net.AddressFamilyIPv4: + specified = true + ipBytes = localAddress.Address.IP().To4() + case net.AddressFamilyIPv6: + specified = true + ipBytes = localAddress.Address.IP().To16() + case net.AddressFamilyDomain: + // Listening on a domain name is not supported. + return nil, fmt.Errorf("listening on domain address not supported: %s", localAddress.Address.String()) + default: + // If unspecified (zero) address, allow kernel (stack) to choose. + specified = false + } + } + + var laddr *tcpip.FullAddress + if specified { + if ipBytes == nil { + return nil, fmt.Errorf("failed to obtain IP bytes for %v", localAddress.Address) + } + netProto := ipv4.ProtocolNumber + if len(ipBytes) == 16 { + netProto = ipv6.ProtocolNumber + } + l := tcpip.FullAddress{Addr: tcpip.AddrFromSlice(ipBytes), Port: uint16(localAddress.Port)} + laddr = &l + // Create UDP endpoint bound to local address. + udpConn, err := gonet.DialUDP(w.stack, laddr, nil, netProto) + if err != nil { + return nil, err + } + return udpConn, nil + } + + // If not specified, let the stack choose the local address (pass nil laddr). + // Default network selection honors PreferIpv6ForUdp if configured. + defaultNet := ipv4.ProtocolNumber + if w.config != nil && w.config.GetPreferIpv6ForUdp() { + defaultNet = ipv6.ProtocolNumber + } + udpConn, err := gonet.DialUDP(w.stack, nil, nil, defaultNet) + if err != nil { + return nil, err + } + return udpConn, nil +} diff --git a/common/packetswitch/gvisorstack/stack.go b/common/packetswitch/gvisorstack/stack.go new file mode 100644 index 000000000..90ccdca25 --- /dev/null +++ b/common/packetswitch/gvisorstack/stack.go @@ -0,0 +1,164 @@ +package gvisorstack + +import ( + "context" + "fmt" + + "github.com/v2fly/v2ray-core/v5/common/packetswitch" + "gvisor.dev/gvisor/pkg/tcpip" + "gvisor.dev/gvisor/pkg/tcpip/network/ipv4" + "gvisor.dev/gvisor/pkg/tcpip/network/ipv6" + "gvisor.dev/gvisor/pkg/tcpip/stack" + "gvisor.dev/gvisor/pkg/tcpip/transport/icmp" + "gvisor.dev/gvisor/pkg/tcpip/transport/tcp" + "gvisor.dev/gvisor/pkg/tcpip/transport/udp" +) + +//go:generate go run github.com/v2fly/v2ray-core/v5/common/errors/errorgen + +type WrappedStack struct { + config *Config + ctx context.Context + stack *stack.Stack +} + +func NewStack(ctx context.Context, config *Config) (*WrappedStack, error) { + return &WrappedStack{ + config: config, + ctx: ctx, + }, nil +} + +func (w *WrappedStack) CreateStackFromNetworkLayerDevice(packetSwitchDevice packetswitch.NetworkLayerDevice) error { + // Validate + if w == nil || w.config == nil { + return fmt.Errorf("no config") + } + + // Determine MTU from config (0 means unspecified) + mtu := int(w.config.GetMtu()) + + // Create adaptor that implements stack.LinkEndpoint + adaptor := NewNetworkLayerDeviceToGvisorLinkEndpointAdaptor(w.ctx, mtu, packetSwitchDevice) + + // Create stack using adaptor as link endpoint + s, err := w.createStack(adaptor) + if err != nil { + // cleanup adaptor on error + adaptor.Close() + return fmt.Errorf("failed to create gvisor stack: %v", err) + } + + // When the adaptor is closed, close the stack as well. + adaptor.SetOnCloseAction(func() { + if s != nil { + s.Close() + } + }) + + w.stack = s + return nil +} + +func (w *WrappedStack) createStack(linkedEndpoint stack.LinkEndpoint) (*stack.Stack, error) { + // Machine Generated + if w == nil || w.config == nil { + return nil, fmt.Errorf("no config") + } + + s := stack.New(stack.Options{ + NetworkProtocols: []stack.NetworkProtocolFactory{ + ipv4.NewProtocol, + ipv6.NewProtocol, + }, + TransportProtocols: []stack.TransportProtocolFactory{ + tcp.NewProtocol, + udp.NewProtocol, + icmp.NewProtocol4, + icmp.NewProtocol6, + }, + }) + + nicID := s.NextNICID() + + // Create NIC + if err := s.CreateNICWithOptions(nicID, linkedEndpoint, stack.NICOptions{Disabled: false, QDisc: nil}); err != nil { + return nil, fmt.Errorf("failed to create NIC: %v", err) + } + + // Add protocol addresses + for _, ip := range w.config.Ips { + tcpIPAddr := tcpip.AddrFrom4Slice(ip.Ip) + protocolAddress := tcpip.ProtocolAddress{ + AddressWithPrefix: tcpip.AddressWithPrefix{ + Address: tcpIPAddr, + PrefixLen: int(ip.Prefix), + }, + } + + switch tcpIPAddr.Len() { + case 4: + protocolAddress.Protocol = ipv4.ProtocolNumber + case 16: + protocolAddress.Protocol = ipv6.ProtocolNumber + default: + return nil, fmt.Errorf("invalid IP address length: %d", tcpIPAddr.Len()) + } + + if err := s.AddProtocolAddress(nicID, protocolAddress, stack.AddressProperties{}); err != nil { + return nil, fmt.Errorf("failed to add protocol address: %v", err) + } + } + + // Set route table + s.SetRouteTable(func() (table []tcpip.Route) { + for _, cidrs := range w.config.Routes { + subnet := tcpip.AddressWithPrefix{ + Address: tcpip.AddrFrom4Slice(cidrs.Ip), + PrefixLen: int(cidrs.Prefix), + }.Subnet() + route := tcpip.Route{ + Destination: subnet, + NIC: nicID, + } + table = append(table, route) + } + return + }()) + + // Promiscuous & spoofing + if err := s.SetPromiscuousMode(nicID, w.config.EnablePromiscuousMode); err != nil { + return nil, fmt.Errorf("failed to set promiscuous mode: %v", err) + } + if err := s.SetSpoofing(nicID, w.config.EnableSpoofing); err != nil { + return nil, fmt.Errorf("failed to set spoofing: %v", err) + } + + // Apply socket buffer sizes if provided + if w.config.SocketSettings != nil { + if size := w.config.SocketSettings.TxBufSize; size != 0 { + sendBufferSizeRangeOption := tcpip.TCPSendBufferSizeRangeOption{Min: tcp.MinBufferSize, Default: int(size), Max: tcp.MaxBufferSize} + if err := s.SetTransportProtocolOption(tcp.ProtocolNumber, &sendBufferSizeRangeOption); err != nil { + return nil, fmt.Errorf("failed to set tcp send buffer size: %v", err) + } + } + + if size := w.config.SocketSettings.RxBufSize; size != 0 { + receiveBufferSizeRangeOption := tcpip.TCPReceiveBufferSizeRangeOption{Min: tcp.MinBufferSize, Default: int(size), Max: tcp.MaxBufferSize} + if err := s.SetTransportProtocolOption(tcp.ProtocolNumber, &receiveBufferSizeRangeOption); err != nil { + return nil, fmt.Errorf("failed to set tcp receive buffer size: %v", err) + } + } + } + + return s, nil +} + +func (w *WrappedStack) Close() error { + if w == nil || w.stack == nil { + return nil + } + w.stack.Close() + w.stack = nil + return nil +} diff --git a/common/packetswitch/interconnect/interconnect.go b/common/packetswitch/interconnect/interconnect.go new file mode 100644 index 000000000..f76d68116 --- /dev/null +++ b/common/packetswitch/interconnect/interconnect.go @@ -0,0 +1,3 @@ +package interconnect + +//go:generate go run github.com/v2fly/v2ray-core/v5/common/errors/errorgen diff --git a/common/packetswitch/interconnect/networkLayer_cable.go b/common/packetswitch/interconnect/networkLayer_cable.go new file mode 100644 index 000000000..3a2f0c0ce --- /dev/null +++ b/common/packetswitch/interconnect/networkLayer_cable.go @@ -0,0 +1,92 @@ +package interconnect + +import ( + "context" + "errors" + "sync" + + "github.com/v2fly/v2ray-core/v5/common/packetswitch" +) + +func NewNetworkLayerCable(ctx context.Context) (*NetworkLayerCable, error) { + return &NetworkLayerCable{ + ctx: ctx, + }, nil +} + +// NetworkLayerCable is primarily Machine Generated +type NetworkLayerCable struct { + lSideWriter packetswitch.NetworkLayerPacketWriter + rSideWriter packetswitch.NetworkLayerPacketWriter + ctx context.Context + lock sync.RWMutex +} + +// NetworkLayerCableDevice is Machine Generated +type NetworkLayerCableDevice struct { + cable *NetworkLayerCable + isLeft bool +} + +func (c *NetworkLayerCable) GetLSideDevice() *NetworkLayerCableDevice { + return &NetworkLayerCableDevice{ + cable: c, + isLeft: true, + } +} + +func (c *NetworkLayerCable) GetRSideDevice() *NetworkLayerCableDevice { + return &NetworkLayerCableDevice{ + cable: c, + isLeft: false, + } +} + +// OnAttach implements NetworkLayerPacketReader.OnAttach +func (d *NetworkLayerCableDevice) OnAttach(writer packetswitch.NetworkLayerPacketWriter) error { + if writer == nil { + return errors.New("nil writer") + } + d.cable.lock.Lock() + defer d.cable.lock.Unlock() + if d.isLeft { + if d.cable.lSideWriter != nil { + return errors.New("left writer already attached") + } + d.cable.lSideWriter = writer + } else { + if d.cable.rSideWriter != nil { + return errors.New("right writer already attached") + } + d.cable.rSideWriter = writer + } + return nil +} + +// Write implements NetworkLayerPacketWriter.Write +func (d *NetworkLayerCableDevice) Write(packet []byte) (int, error) { + d.cable.lock.RLock() + var peer packetswitch.NetworkLayerPacketWriter + if d.isLeft { + peer = d.cable.rSideWriter + } else { + peer = d.cable.lSideWriter + } + d.cable.lock.RUnlock() + if peer == nil { + return 0, errors.New("no peer attached") + } + return peer.Write(packet) +} + +// Close implements common.Closable.Close +func (d *NetworkLayerCableDevice) Close() error { + d.cable.lock.Lock() + defer d.cable.lock.Unlock() + if d.isLeft { + d.cable.lSideWriter = nil + } else { + d.cable.rSideWriter = nil + } + return nil +} diff --git a/common/packetswitch/interconnect/networkLayer_cable_test.go b/common/packetswitch/interconnect/networkLayer_cable_test.go new file mode 100644 index 000000000..d8a6fad07 --- /dev/null +++ b/common/packetswitch/interconnect/networkLayer_cable_test.go @@ -0,0 +1,182 @@ +package interconnect + +import ( + "context" + "fmt" + "sync" + "testing" + "time" + + "github.com/v2fly/v2ray-core/v5/common/packetswitch" +) + +type testWriter struct { + mu sync.Mutex + received [][]byte + ch chan []byte +} + +func newTestWriter(bufSize int) *testWriter { + w := &testWriter{ch: make(chan []byte, bufSize)} + return w +} + +func (w *testWriter) Write(p []byte) (int, error) { + // copy payload + cp := make([]byte, len(p)) + copy(cp, p) + w.mu.Lock() + w.received = append(w.received, cp) + w.mu.Unlock() + select { + case w.ch <- cp: + default: + } + return len(p), nil +} + +func (w *testWriter) ReceivedAll() [][]byte { + w.mu.Lock() + defer w.mu.Unlock() + out := make([][]byte, len(w.received)) + copy(out, w.received) + return out +} + +func TestCable_HappyPath(t *testing.T) { + c, err := NewNetworkLayerCable(context.Background()) + if err != nil { + t.Fatalf("failed to create cable: %v", err) + } + l := c.GetLSideDevice() + r := c.GetRSideDevice() + + wL := newTestWriter(4) + wR := newTestWriter(4) + + if err := l.OnAttach(wL); err != nil { + t.Fatalf("attach left failed: %v", err) + } + if err := r.OnAttach(wR); err != nil { + t.Fatalf("attach right failed: %v", err) + } + + payloadL := []byte("from-left") + n, err := l.Write(payloadL) + if err != nil { + t.Fatalf("write left failed: %v", err) + } + if n != len(payloadL) { + t.Fatalf("write returned wrong length: %d", n) + } + + select { + case got := <-wR.ch: + if string(got) != string(payloadL) { + t.Fatalf("unexpected payload at right: %s", string(got)) + } + case <-time.After(500 * time.Millisecond): + t.Fatalf("timeout waiting for payload on right") + } + + payloadR := []byte("from-right") + n, err = r.Write(payloadR) + if err != nil { + t.Fatalf("write right failed: %v", err) + } + if n != len(payloadR) { + t.Fatalf("write returned wrong length: %d", n) + } + select { + case got := <-wL.ch: + if string(got) != string(payloadR) { + t.Fatalf("unexpected payload at left: %s", string(got)) + } + case <-time.After(500 * time.Millisecond): + t.Fatalf("timeout waiting for payload on left") + } +} + +func TestCable_NoPeer(t *testing.T) { + c, _ := NewNetworkLayerCable(context.Background()) + l := c.GetLSideDevice() + wL := newTestWriter(1) + if err := l.OnAttach(wL); err != nil { + t.Fatalf("attach left failed: %v", err) + } + if n, err := l.Write([]byte("x")); err == nil || n != 0 { + t.Fatalf("expected write to fail when no peer attached got n=%d err=%v", n, err) + } +} + +func TestCable_DoubleAttachAndClose(t *testing.T) { + c, _ := NewNetworkLayerCable(context.Background()) + l := c.GetLSideDevice() + w1 := newTestWriter(1) + w2 := newTestWriter(1) + if err := l.OnAttach(w1); err != nil { + t.Fatalf("attach left failed: %v", err) + } + if err := l.OnAttach(w2); err == nil { + t.Fatalf("expected second attach to fail") + } + + r := c.GetRSideDevice() + wr := newTestWriter(2) + if err := r.OnAttach(wr); err != nil { + t.Fatalf("attach right failed: %v", err) + } + + // close left and ensure right cannot write + if err := l.Close(); err != nil { + t.Fatalf("close left failed: %v", err) + } + if n, err := r.Write([]byte("hello")); err == nil || n != 0 { + t.Fatalf("expected write from right to fail after left closed got n=%d err=%v", n, err) + } +} + +func TestCable_ConcurrentWrites(t *testing.T) { + c, _ := NewNetworkLayerCable(context.Background()) + l := c.GetLSideDevice() + r := c.GetRSideDevice() + wL := newTestWriter(100) + wR := newTestWriter(100) + if err := l.OnAttach(wL); err != nil { + t.Fatalf("attach left failed: %v", err) + } + if err := r.OnAttach(wR); err != nil { + t.Fatalf("attach right failed: %v", err) + } + + var wg sync.WaitGroup + count := 200 + wg.Add(count * 2) + for i := 0; i < count; i++ { + payloadL := []byte(fmt.Sprintf("L-%d", i%10)) + payloadR := []byte(fmt.Sprintf("R-%d", i%10)) + go func(p []byte) { + defer wg.Done() + _, _ = l.Write(p) + }(payloadL) + go func(p []byte) { + defer wg.Done() + _, _ = r.Write(p) + }(payloadR) + } + wg.Wait() + + // drain channels (best-effort) + timed := time.After(500 * time.Millisecond) + for { + select { + case <-wL.ch: + case <-wR.ch: + case <-timed: + return + } + } +} + +// ensure testWriter implements the interface +var _ packetswitch.NetworkLayerPacketWriter = (*testWriter)(nil) diff --git a/common/packetswitch/packetswitch.go b/common/packetswitch/packetswitch.go new file mode 100644 index 000000000..83e3b2dad --- /dev/null +++ b/common/packetswitch/packetswitch.go @@ -0,0 +1,17 @@ +package packetswitch + +import "github.com/v2fly/v2ray-core/v5/common" + +type NetworkLayerDevice interface { + common.Closable + NetworkLayerPacketWriter + NetworkLayerPacketReader +} + +type NetworkLayerPacketWriter interface { + Write(packet []byte) (n int, err error) +} + +type NetworkLayerPacketReader interface { + OnAttach(writer NetworkLayerPacketWriter) error +} diff --git a/go.mod b/go.mod index bbaa88b8e..b19989732 100644 --- a/go.mod +++ b/go.mod @@ -91,6 +91,8 @@ require ( golang.org/x/text v0.34.0 // indirect golang.org/x/time v0.12.0 // indirect golang.org/x/tools v0.41.0 // indirect + golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect + golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb // indirect google.golang.org/genproto/googleapis/rpc v0.0.0-20251202230838-ff82c1b0f217 // indirect nhooyr.io/websocket v1.8.6 // indirect ) diff --git a/go.sum b/go.sum index ac193a68f..037c31b94 100644 --- a/go.sum +++ b/go.sum @@ -788,6 +788,10 @@ golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8T golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= +golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg= +golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI= +golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A= +golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw= gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk= gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E= google.golang.org/api v0.3.1/go.mod h1:6wY9I6uQWHQ8EM57III9mq/AjF+i8G65rmVagqKMtkk= diff --git a/main/distro/all/all.go b/main/distro/all/all.go index 8476f957d..9e0c44c57 100644 --- a/main/distro/all/all.go +++ b/main/distro/all/all.go @@ -58,6 +58,9 @@ import ( _ "github.com/v2fly/v2ray-core/v5/proxy/hysteria2" _ "github.com/v2fly/v2ray-core/v5/proxy/shadowsocks2022" + // WireGuard Outbound is unreleased. + _ "github.com/v2fly/v2ray-core/v5/proxy/wireguard/outbound" + // Transports _ "github.com/v2fly/v2ray-core/v5/transport/internet/domainsocket" _ "github.com/v2fly/v2ray-core/v5/transport/internet/grpc" diff --git a/proxy/wireguard/outbound/config.pb.go b/proxy/wireguard/outbound/config.pb.go new file mode 100644 index 000000000..3385f975c --- /dev/null +++ b/proxy/wireguard/outbound/config.pb.go @@ -0,0 +1,216 @@ +package outbound + +import ( + _ "github.com/v2fly/v2ray-core/v5/common/net/packetaddr" + gvisorstack "github.com/v2fly/v2ray-core/v5/common/packetswitch/gvisorstack" + _ "github.com/v2fly/v2ray-core/v5/common/protoext" + wgcommon "github.com/v2fly/v2ray-core/v5/proxy/wireguard/wgcommon" + protoreflect "google.golang.org/protobuf/reflect/protoreflect" + protoimpl "google.golang.org/protobuf/runtime/protoimpl" + reflect "reflect" + sync "sync" + unsafe "unsafe" +) + +const ( + // Verify that this generated code is sufficiently up-to-date. + _ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion) + // Verify that runtime/protoimpl is sufficiently up-to-date. + _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) +) + +type Config_DomainStrategy int32 + +const ( + Config_AS_IS Config_DomainStrategy = 0 + Config_USE_IP Config_DomainStrategy = 1 + Config_USE_IP4 Config_DomainStrategy = 2 + Config_USE_IP6 Config_DomainStrategy = 3 +) + +// Enum value maps for Config_DomainStrategy. +var ( + Config_DomainStrategy_name = map[int32]string{ + 0: "AS_IS", + 1: "USE_IP", + 2: "USE_IP4", + 3: "USE_IP6", + } + Config_DomainStrategy_value = map[string]int32{ + "AS_IS": 0, + "USE_IP": 1, + "USE_IP4": 2, + "USE_IP6": 3, + } +) + +func (x Config_DomainStrategy) Enum() *Config_DomainStrategy { + p := new(Config_DomainStrategy) + *p = x + return p +} + +func (x Config_DomainStrategy) String() string { + return protoimpl.X.EnumStringOf(x.Descriptor(), protoreflect.EnumNumber(x)) +} + +func (Config_DomainStrategy) Descriptor() protoreflect.EnumDescriptor { + return file_proxy_wireguard_outbound_config_proto_enumTypes[0].Descriptor() +} + +func (Config_DomainStrategy) Type() protoreflect.EnumType { + return &file_proxy_wireguard_outbound_config_proto_enumTypes[0] +} + +func (x Config_DomainStrategy) Number() protoreflect.EnumNumber { + return protoreflect.EnumNumber(x) +} + +// Deprecated: Use Config_DomainStrategy.Descriptor instead. +func (Config_DomainStrategy) EnumDescriptor() ([]byte, []int) { + return file_proxy_wireguard_outbound_config_proto_rawDescGZIP(), []int{0, 0} +} + +type Config struct { + state protoimpl.MessageState `protogen:"open.v1"` + WgDevice *wgcommon.DeviceConfig `protobuf:"bytes,1,opt,name=wg_device,json=wgDevice,proto3" json:"wg_device,omitempty"` + Stack *gvisorstack.Config `protobuf:"bytes,2,opt,name=stack,proto3" json:"stack,omitempty"` + // v2ray.core.net.packetaddr.PacketAddrType outbound_packet_encoding = 3; + ListenOnSystemNetwork bool `protobuf:"varint,4,opt,name=listen_on_system_network,json=listenOnSystemNetwork,proto3" json:"listen_on_system_network,omitempty"` + DomainStrategy Config_DomainStrategy `protobuf:"varint,5,opt,name=domain_strategy,json=domainStrategy,proto3,enum=v2ray.core.proxy.wireguard.outbound.Config_DomainStrategy" json:"domain_strategy,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *Config) Reset() { + *x = Config{} + mi := &file_proxy_wireguard_outbound_config_proto_msgTypes[0] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *Config) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*Config) ProtoMessage() {} + +func (x *Config) ProtoReflect() protoreflect.Message { + mi := &file_proxy_wireguard_outbound_config_proto_msgTypes[0] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use Config.ProtoReflect.Descriptor instead. +func (*Config) Descriptor() ([]byte, []int) { + return file_proxy_wireguard_outbound_config_proto_rawDescGZIP(), []int{0} +} + +func (x *Config) GetWgDevice() *wgcommon.DeviceConfig { + if x != nil { + return x.WgDevice + } + return nil +} + +func (x *Config) GetStack() *gvisorstack.Config { + if x != nil { + return x.Stack + } + return nil +} + +func (x *Config) GetListenOnSystemNetwork() bool { + if x != nil { + return x.ListenOnSystemNetwork + } + return false +} + +func (x *Config) GetDomainStrategy() Config_DomainStrategy { + if x != nil { + return x.DomainStrategy + } + return Config_AS_IS +} + +var File_proxy_wireguard_outbound_config_proto protoreflect.FileDescriptor + +const file_proxy_wireguard_outbound_config_proto_rawDesc = "" + + "\n" + + "%proxy/wireguard/outbound/config.proto\x12#v2ray.core.proxy.wireguard.outbound\x1a%proxy/wireguard/wgcommon/config.proto\x1a,common/packetswitch/gvisorstack/config.proto\x1a\"common/net/packetaddr/config.proto\x1a common/protoext/extensions.proto\"\x9e\x03\n" + + "\x06Config\x12N\n" + + "\twg_device\x18\x01 \x01(\v21.v2ray.core.proxy.wireguard.wgcommon.DeviceConfigR\bwgDevice\x12H\n" + + "\x05stack\x18\x02 \x01(\v22.v2ray.core.common.packetswitch.gvisorstack.ConfigR\x05stack\x127\n" + + "\x18listen_on_system_network\x18\x04 \x01(\bR\x15listenOnSystemNetwork\x12c\n" + + "\x0fdomain_strategy\x18\x05 \x01(\x0e2:.v2ray.core.proxy.wireguard.outbound.Config.DomainStrategyR\x0edomainStrategy\"A\n" + + "\x0eDomainStrategy\x12\t\n" + + "\x05AS_IS\x10\x00\x12\n" + + "\n" + + "\x06USE_IP\x10\x01\x12\v\n" + + "\aUSE_IP4\x10\x02\x12\v\n" + + "\aUSE_IP6\x10\x03:\x19\x82\xb5\x18\x15\n" + + "\boutbound\x12\twireguardB\x8a\x01\n" + + "'com.v2ray.core.proxy.wireguard.outboundP\x01Z7github.com/v2fly/v2ray-core/v5/proxy/wireguard/outbound\xaa\x02#V2Ray.Core.Proxy.Wireguard.Outboundb\x06proto3" + +var ( + file_proxy_wireguard_outbound_config_proto_rawDescOnce sync.Once + file_proxy_wireguard_outbound_config_proto_rawDescData []byte +) + +func file_proxy_wireguard_outbound_config_proto_rawDescGZIP() []byte { + file_proxy_wireguard_outbound_config_proto_rawDescOnce.Do(func() { + file_proxy_wireguard_outbound_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_proxy_wireguard_outbound_config_proto_rawDesc), len(file_proxy_wireguard_outbound_config_proto_rawDesc))) + }) + return file_proxy_wireguard_outbound_config_proto_rawDescData +} + +var file_proxy_wireguard_outbound_config_proto_enumTypes = make([]protoimpl.EnumInfo, 1) +var file_proxy_wireguard_outbound_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1) +var file_proxy_wireguard_outbound_config_proto_goTypes = []any{ + (Config_DomainStrategy)(0), // 0: v2ray.core.proxy.wireguard.outbound.Config.DomainStrategy + (*Config)(nil), // 1: v2ray.core.proxy.wireguard.outbound.Config + (*wgcommon.DeviceConfig)(nil), // 2: v2ray.core.proxy.wireguard.wgcommon.DeviceConfig + (*gvisorstack.Config)(nil), // 3: v2ray.core.common.packetswitch.gvisorstack.Config +} +var file_proxy_wireguard_outbound_config_proto_depIdxs = []int32{ + 2, // 0: v2ray.core.proxy.wireguard.outbound.Config.wg_device:type_name -> v2ray.core.proxy.wireguard.wgcommon.DeviceConfig + 3, // 1: v2ray.core.proxy.wireguard.outbound.Config.stack:type_name -> v2ray.core.common.packetswitch.gvisorstack.Config + 0, // 2: v2ray.core.proxy.wireguard.outbound.Config.domain_strategy:type_name -> v2ray.core.proxy.wireguard.outbound.Config.DomainStrategy + 3, // [3:3] is the sub-list for method output_type + 3, // [3:3] is the sub-list for method input_type + 3, // [3:3] is the sub-list for extension type_name + 3, // [3:3] is the sub-list for extension extendee + 0, // [0:3] is the sub-list for field type_name +} + +func init() { file_proxy_wireguard_outbound_config_proto_init() } +func file_proxy_wireguard_outbound_config_proto_init() { + if File_proxy_wireguard_outbound_config_proto != nil { + return + } + type x struct{} + out := protoimpl.TypeBuilder{ + File: protoimpl.DescBuilder{ + GoPackagePath: reflect.TypeOf(x{}).PkgPath(), + RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_wireguard_outbound_config_proto_rawDesc), len(file_proxy_wireguard_outbound_config_proto_rawDesc)), + NumEnums: 1, + NumMessages: 1, + NumExtensions: 0, + NumServices: 0, + }, + GoTypes: file_proxy_wireguard_outbound_config_proto_goTypes, + DependencyIndexes: file_proxy_wireguard_outbound_config_proto_depIdxs, + EnumInfos: file_proxy_wireguard_outbound_config_proto_enumTypes, + MessageInfos: file_proxy_wireguard_outbound_config_proto_msgTypes, + }.Build() + File_proxy_wireguard_outbound_config_proto = out.File + file_proxy_wireguard_outbound_config_proto_goTypes = nil + file_proxy_wireguard_outbound_config_proto_depIdxs = nil +} diff --git a/proxy/wireguard/outbound/config.proto b/proxy/wireguard/outbound/config.proto new file mode 100644 index 000000000..7fe226c45 --- /dev/null +++ b/proxy/wireguard/outbound/config.proto @@ -0,0 +1,31 @@ +syntax = "proto3"; + +package v2ray.core.proxy.wireguard.outbound; +option csharp_namespace = "V2Ray.Core.Proxy.Wireguard.Outbound"; +option go_package = "github.com/v2fly/v2ray-core/v5/proxy/wireguard/outbound"; +option java_package = "com.v2ray.core.proxy.wireguard.outbound"; +option java_multiple_files = true; + +import "proxy/wireguard/wgcommon/config.proto"; +import "common/packetswitch/gvisorstack/config.proto"; +import "common/net/packetaddr/config.proto"; +import "common/protoext/extensions.proto"; + +message Config{ + option (v2ray.core.common.protoext.message_opt).type = "outbound"; + option (v2ray.core.common.protoext.message_opt).short_name = "wireguard"; + + v2ray.core.proxy.wireguard.wgcommon.DeviceConfig wg_device = 1; + v2ray.core.common.packetswitch.gvisorstack.Config stack = 2; + + //v2ray.core.net.packetaddr.PacketAddrType outbound_packet_encoding = 3; + bool listen_on_system_network = 4; + + enum DomainStrategy { + AS_IS = 0; + USE_IP = 1; + USE_IP4 = 2; + USE_IP6 = 3; + } + DomainStrategy domain_strategy = 5; +} \ No newline at end of file diff --git a/proxy/wireguard/outbound/errors.generated.go b/proxy/wireguard/outbound/errors.generated.go new file mode 100644 index 000000000..44135e859 --- /dev/null +++ b/proxy/wireguard/outbound/errors.generated.go @@ -0,0 +1,9 @@ +package outbound + +import "github.com/v2fly/v2ray-core/v5/common/errors" + +type errPathObjHolder struct{} + +func newError(values ...interface{}) *errors.Error { + return errors.New(values...).WithPathObj(errPathObjHolder{}) +} diff --git a/proxy/wireguard/outbound/outbound.go b/proxy/wireguard/outbound/outbound.go new file mode 100644 index 000000000..914415af4 --- /dev/null +++ b/proxy/wireguard/outbound/outbound.go @@ -0,0 +1,387 @@ +package outbound + +import ( + "context" + gonet "net" + "sync" + "time" + + core "github.com/v2fly/v2ray-core/v5" + "github.com/v2fly/v2ray-core/v5/common" + "github.com/v2fly/v2ray-core/v5/common/buf" + "github.com/v2fly/v2ray-core/v5/common/dice" + "github.com/v2fly/v2ray-core/v5/common/environment" + "github.com/v2fly/v2ray-core/v5/common/environment/envctx" + cnet "github.com/v2fly/v2ray-core/v5/common/net" + "github.com/v2fly/v2ray-core/v5/common/net/packetaddr" + "github.com/v2fly/v2ray-core/v5/common/packetswitch/gvisorstack" + "github.com/v2fly/v2ray-core/v5/common/packetswitch/interconnect" + "github.com/v2fly/v2ray-core/v5/common/session" + "github.com/v2fly/v2ray-core/v5/common/signal" + "github.com/v2fly/v2ray-core/v5/common/task" + "github.com/v2fly/v2ray-core/v5/features/dns" + "github.com/v2fly/v2ray-core/v5/proxy/wireguard/wgcommon" + "github.com/v2fly/v2ray-core/v5/transport" + "github.com/v2fly/v2ray-core/v5/transport/internet" + "github.com/v2fly/v2ray-core/v5/transport/internet/udp" +) + +//go:generate go run github.com/v2fly/v2ray-core/v5/common/errors/errorgen + +func NewWireguardOutbound(ctx context.Context, config *Config) (*WireguardOutbound, error) { + w := &WireguardOutbound{ + ctx: ctx, + config: config, + } + // Acquire dns client feature if available + if err := core.RequireFeatures(ctx, func(d dns.Client) error { + w.dnsClient = d + return nil + }); err != nil { + return nil, newError("failed to require dns client feature").Base(err) + } + storage := envctx.EnvironmentFromContext(ctx).(environment.ProxyEnvironment).TransientStorage() + + udpState, err := NewClientConnState() + if err != nil { + return nil, newError("failed to create UDP connection state").Base(err) + } + if err := storage.Put(ctx, ConnectionState, udpState); err != nil { + return nil, newError("failed to put connection state").Base(err) + } + return w, nil +} + +type WireguardOutbound struct { + ctx context.Context + config *Config + + dnsClient dns.Client +} + +type WireguardOutboundSession struct { + ctx context.Context + config *Config + + stack *gvisorstack.WrappedStack + wireguardDevice *wgcommon.WrappedWireguardDevice + interconnect *interconnect.NetworkLayerCable + + // system packet conn used when ListenOnSystemNetwork is true + systemPacketConn internet.PacketConn + + dnsClient dns.Client +} + +func (s *WireguardOutboundSession) initFromConfig(ctx context.Context, config *Config) error { + if config == nil { + return newError("nil config") + } + // create interconnect cable + cable, err := interconnect.NewNetworkLayerCable(ctx) + if err != nil { + return newError("failed to create interconnect cable").Base(err) + } + s.interconnect = cable + + // create wireguard device wrapper + wd, err := wgcommon.NewWrappedWireguardDevice(ctx, config.GetWgDevice()) + if err != nil { + return newError("failed to create wireguard device").Base(err) + } + s.wireguardDevice = wd + // attach device tunnel to left side of cable + s.wireguardDevice.SetTunnel(cable.GetLSideDevice()) + + // create gvisor stack wrapper if stack config is provided + if config.GetStack() != nil { + st, err := gvisorstack.NewStack(ctx, config.GetStack()) + if err != nil { + return newError("failed to create gvisor stack").Base(err) + } + s.stack = st + if err := s.stack.CreateStackFromNetworkLayerDevice(cable.GetRSideDevice()); err != nil { + return newError("failed to create stack from network layer device").Base(err) + } + } + + return nil +} + +const ConnectionState = "ConnectionState" + +type ClientConnState struct { + session *WireguardOutboundSession + initOnce *sync.Once + mu sync.Mutex +} + +func (c *ClientConnState) GetOrCreateSession(create func() (*WireguardOutboundSession, error)) (*WireguardOutboundSession, error) { + var errOuter error + c.initOnce.Do(func() { + sess, err := create() + if err != nil { + errOuter = err + return + } + c.mu.Lock() + c.session = sess + c.mu.Unlock() + }) + if errOuter != nil { + return nil, newError("failed to initialize UDP State").Base(errOuter) + } + return c.session, nil +} + +func (c *ClientConnState) IsTransientStorageLifecycleReceiver() {} + +func (c *ClientConnState) Close() error { + c.mu.Lock() + defer c.mu.Unlock() + if c.session == nil { + return nil + } + sess := c.session + c.session = nil + + // close interconnect devices first to stop any further packet injections + if sess.interconnect != nil { + _ = sess.interconnect.GetLSideDevice().Close() + _ = sess.interconnect.GetRSideDevice().Close() + sess.interconnect = nil + } + + // close system packet conn + if sess.systemPacketConn != nil { + _ = sess.systemPacketConn.Close() + sess.systemPacketConn = nil + } + + // close wireguard device + if sess.wireguardDevice != nil { + _ = sess.wireguardDevice.Close() + sess.wireguardDevice = nil + } + + // Close stack last to quiesce any gVisor internal goroutines that may + // hold references to PacketBuffers (prevents dec-ref races). + if sess.stack != nil { + _ = sess.stack.Close() + sess.stack = nil + } + + return nil +} + +func (w *WireguardOutbound) Process(ctx context.Context, link *transport.Link, dialer internet.Dialer) error { + // keep dialer for address family preference when resolving domain + _ = dialer + storage := envctx.EnvironmentFromContext(w.ctx).(environment.ProxyEnvironment).TransientStorage() + stateIfc, err := storage.Get(ctx, ConnectionState) + if err != nil { + return newError("failed to get connection state").Base(err) + } + clientState, ok := stateIfc.(*ClientConnState) + if !ok { + return newError("bad connection state") + } + + // create session if needed + sess, err := clientState.GetOrCreateSession(func() (*WireguardOutboundSession, error) { + s := &WireguardOutboundSession{ctx: ctx, config: w.config} + s.dnsClient = w.dnsClient + if err := s.initFromConfig(ctx, w.config); err != nil { + return nil, err + } + + if !w.config.ListenOnSystemNetwork { + // SORRRRRY, I tried but it was v2ray's udp support was too difficult to work with + return nil, newError("unimplemented: listenOnSystemNetwork=false is not implemented yet") + } + + packetConn, err := internet.ListenSystemPacket(w.ctx, &gonet.UDPAddr{IP: cnet.AnyIP.IP(), Port: 0}, nil) + if err != nil { + return nil, newError("failed to listen on system network").Base(err) + } + + s.systemPacketConn = packetConn + s.wireguardDevice.SetConn(packetConn) + + // initialize wireguard device now that conn present + if err := s.wireguardDevice.InitDevice(); err != nil { + return nil, newError("failed to init wireguard device").Base(err) + } + if err := s.wireguardDevice.SetupDeviceWithoutPeers(); err != nil { + return nil, newError("failed to setup wireguard device").Base(err) + } + if err := s.wireguardDevice.AddOrReplacePeers(s.config.WgDevice.GetPeers()); err != nil { + return nil, newError("failed to add peers").Base(err) + } + if err := s.wireguardDevice.Up(); err != nil { + return nil, newError("failed to bring up wireguard device").Base(err) + } + return s, nil + }) + if err != nil { + return newError("failed to create or fetch session").Base(err) + } + + { + debugData, err := sess.wireguardDevice.Debug() + if err != nil { + newError("failed to debug wireguard device").Base(err).WriteToLog(session.ExportIDToError(ctx)) + } + newError("wireguard device debug: \n", debugData).AtDebug().WriteToLog(session.ExportIDToError(ctx)) + } + + outbound := session.OutboundFromContext(ctx) + if outbound == nil || !outbound.Target.IsValid() { + return newError("target not specified") + } + destination := outbound.Target + + // resolve domain names using dns client if necessary + if destination.Address != nil && destination.Address.Family().IsDomain() && sess.dnsClient != nil { + domain := destination.Address.Domain() + // determine IP option based on domain strategy and dialer local address + var localAddr cnet.Address + if dialer != nil { + localAddr = dialer.Address() + } + opt := dns.IPOption{ + IPv4Enable: sess.config.DomainStrategy == Config_USE_IP || sess.config.DomainStrategy == Config_USE_IP4 || (localAddr != nil && localAddr.Family().IsIPv4()), + IPv6Enable: sess.config.DomainStrategy == Config_USE_IP || sess.config.DomainStrategy == Config_USE_IP6 || (localAddr != nil && localAddr.Family().IsIPv6()), + FakeEnable: false, + } + ips, err := dns.LookupIPWithOption(sess.dnsClient, domain, opt) + if err != nil { + newError("failed to get IP address for domain ", domain).Base(err).WriteToLog(session.ExportIDToError(ctx)) + } + if len(ips) > 0 { + // pick a random IP from results + ip := ips[dice.Roll(len(ips))] + destination.Address = cnet.IPAddress(ip) + } else { + // no resolved IP; continue and let later operations fail appropriately + newError("no IP resolved for domain ", domain).AtWarning().WriteToLog(session.ExportIDToError(ctx)) + } + } + + // require gVisor stack to process network-level connections + if sess.stack == nil { + return newError("gvisor stack is not configured for wireguard outbound") + } + + ctx, cancel := context.WithCancel(ctx) + timer := signal.CancelAfterInactivity(ctx, cancel, time.Second*300) + defer cancel() + + if packetConn, err := packetaddr.ToPacketAddrConn(link, destination); err == nil { + defer func() { _ = packetConn.Close() }() + pc, err := sess.stack.ListenUDP(ctx, cnet.UDPDestination(cnet.AnyIP, 0)) + if err != nil { + return newError("failed to create udp session in stack").Base(err) + } + defer func() { _ = pc.Close() }() + + // Run copy loops and explicitly close resources afterwards to avoid leaks. + err = nil + func() { + requestDone := func() error { + protocolWriter := pc + return udp.CopyPacketConn(protocolWriter, packetConn, udp.UpdateActivity(timer)) + } + responseDone := func() error { + protocolReader := pc + return udp.CopyPacketConn(packetConn, protocolReader, udp.UpdateActivity(timer)) + } + responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer)) + err = task.Run(ctx, requestDone, responseDoneAndCloseWriter) + }() + + if err != nil { + return newError("connection ends").Base(err) + } + return nil + } + + switch destination.Network { + case cnet.Network_TCP: + // Dial TCP inside the virtual stack + conn, err := sess.stack.DialTCP(ctx, destination) + if err != nil { + return newError("failed to dial tcp in stack").Base(err) + } + defer func() { _ = conn.Close() }() + + requestDone := func() error { + writer := buf.NewWriter(conn) + if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil { + return newError("failed to copy request").Base(err) + } + return nil + } + + responseDone := func() error { + reader := buf.NewReader(conn) + if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil { + return newError("failed to copy response").Base(err) + } + return nil + } + + if err := task.Run(ctx, requestDone, task.OnSuccess(responseDone, task.Close(link.Writer))); err != nil { + return newError("connection ends").Base(err) + } + return nil + + case cnet.Network_UDP: + // Create a packet conn on the stack and use mono-dest adapter + pc, err := sess.stack.ListenUDP(ctx, cnet.UDPDestination(nil, 0)) + if err != nil { + return newError("failed to create udp session in stack").Base(err) + } + mono := udp.NewMonoDestUDPConn(pc, &gonet.UDPAddr{IP: destination.Address.IP(), Port: int(destination.Port)}) + + requestDone := func() error { + return buf.Copy(link.Reader, mono, buf.UpdateActivity(timer)) + } + responseDone := func() error { + return buf.Copy(mono, link.Writer, buf.UpdateActivity(timer)) + } + + if err := task.Run(ctx, requestDone, task.OnSuccess(responseDone, task.Close(link.Writer))); err != nil { + _ = pc.Close() + return newError("connection ends").Base(err) + } + return nil + + default: + return newError("unsupported network: ", destination.Network) + } +} + +func (w *WireguardOutbound) Close() error { + storage := envctx.EnvironmentFromContext(w.ctx).(environment.ProxyEnvironment).TransientStorage() + stateIfc, err := storage.Get(context.Background(), ConnectionState) + if err != nil || stateIfc == nil { + return nil + } + clientState, ok := stateIfc.(*ClientConnState) + if !ok || clientState.session == nil { + return nil + } + _ = clientState.Close() + return nil +} + +func NewClientConnState() (*ClientConnState, error) { + return &ClientConnState{initOnce: &sync.Once{}}, nil +} + +func init() { + common.Must(common.RegisterConfig((*Config)(nil), func(ctx context.Context, config interface{}) (interface{}, error) { + return NewWireguardOutbound(ctx, config.(*Config)) + })) +} diff --git a/proxy/wireguard/wgcommon/config.pb.go b/proxy/wireguard/wgcommon/config.pb.go new file mode 100644 index 000000000..f5e74c1b1 --- /dev/null +++ b/proxy/wireguard/wgcommon/config.pb.go @@ -0,0 +1,233 @@ +package wgcommon + +import ( + protoreflect "google.golang.org/protobuf/reflect/protoreflect" + protoimpl "google.golang.org/protobuf/runtime/protoimpl" + reflect "reflect" + sync "sync" + unsafe "unsafe" +) + +const ( + // Verify that this generated code is sufficiently up-to-date. + _ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion) + // Verify that runtime/protoimpl is sufficiently up-to-date. + _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) +) + +type PeerConfig struct { + state protoimpl.MessageState `protogen:"open.v1"` + PublicKey []byte `protobuf:"bytes,1,opt,name=public_key,json=publicKey,proto3" json:"public_key,omitempty"` + PresharedKey []byte `protobuf:"bytes,2,opt,name=preshared_key,json=presharedKey,proto3" json:"preshared_key,omitempty"` + AllowedIps []string `protobuf:"bytes,3,rep,name=allowed_ips,json=allowedIps,proto3" json:"allowed_ips,omitempty"` + Endpoint string `protobuf:"bytes,4,opt,name=endpoint,proto3" json:"endpoint,omitempty"` + PersistentKeepaliveInterval int64 `protobuf:"varint,5,opt,name=persistent_keepalive_interval,json=persistentKeepaliveInterval,proto3" json:"persistent_keepalive_interval,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *PeerConfig) Reset() { + *x = PeerConfig{} + mi := &file_proxy_wireguard_wgcommon_config_proto_msgTypes[0] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *PeerConfig) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*PeerConfig) ProtoMessage() {} + +func (x *PeerConfig) ProtoReflect() protoreflect.Message { + mi := &file_proxy_wireguard_wgcommon_config_proto_msgTypes[0] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use PeerConfig.ProtoReflect.Descriptor instead. +func (*PeerConfig) Descriptor() ([]byte, []int) { + return file_proxy_wireguard_wgcommon_config_proto_rawDescGZIP(), []int{0} +} + +func (x *PeerConfig) GetPublicKey() []byte { + if x != nil { + return x.PublicKey + } + return nil +} + +func (x *PeerConfig) GetPresharedKey() []byte { + if x != nil { + return x.PresharedKey + } + return nil +} + +func (x *PeerConfig) GetAllowedIps() []string { + if x != nil { + return x.AllowedIps + } + return nil +} + +func (x *PeerConfig) GetEndpoint() string { + if x != nil { + return x.Endpoint + } + return "" +} + +func (x *PeerConfig) GetPersistentKeepaliveInterval() int64 { + if x != nil { + return x.PersistentKeepaliveInterval + } + return 0 +} + +type DeviceConfig struct { + state protoimpl.MessageState `protogen:"open.v1"` + PrivateKey []byte `protobuf:"bytes,1,opt,name=private_key,json=privateKey,proto3" json:"private_key,omitempty"` + ListenPort uint32 `protobuf:"varint,3,opt,name=listen_port,json=listenPort,proto3" json:"listen_port,omitempty"` + Peers []*PeerConfig `protobuf:"bytes,4,rep,name=peers,proto3" json:"peers,omitempty"` + Mtu uint32 `protobuf:"varint,5,opt,name=mtu,proto3" json:"mtu,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *DeviceConfig) Reset() { + *x = DeviceConfig{} + mi := &file_proxy_wireguard_wgcommon_config_proto_msgTypes[1] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *DeviceConfig) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*DeviceConfig) ProtoMessage() {} + +func (x *DeviceConfig) ProtoReflect() protoreflect.Message { + mi := &file_proxy_wireguard_wgcommon_config_proto_msgTypes[1] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use DeviceConfig.ProtoReflect.Descriptor instead. +func (*DeviceConfig) Descriptor() ([]byte, []int) { + return file_proxy_wireguard_wgcommon_config_proto_rawDescGZIP(), []int{1} +} + +func (x *DeviceConfig) GetPrivateKey() []byte { + if x != nil { + return x.PrivateKey + } + return nil +} + +func (x *DeviceConfig) GetListenPort() uint32 { + if x != nil { + return x.ListenPort + } + return 0 +} + +func (x *DeviceConfig) GetPeers() []*PeerConfig { + if x != nil { + return x.Peers + } + return nil +} + +func (x *DeviceConfig) GetMtu() uint32 { + if x != nil { + return x.Mtu + } + return 0 +} + +var File_proxy_wireguard_wgcommon_config_proto protoreflect.FileDescriptor + +const file_proxy_wireguard_wgcommon_config_proto_rawDesc = "" + + "\n" + + "%proxy/wireguard/wgcommon/config.proto\x12#v2ray.core.proxy.wireguard.wgcommon\"\xd1\x01\n" + + "\n" + + "PeerConfig\x12\x1d\n" + + "\n" + + "public_key\x18\x01 \x01(\fR\tpublicKey\x12#\n" + + "\rpreshared_key\x18\x02 \x01(\fR\fpresharedKey\x12\x1f\n" + + "\vallowed_ips\x18\x03 \x03(\tR\n" + + "allowedIps\x12\x1a\n" + + "\bendpoint\x18\x04 \x01(\tR\bendpoint\x12B\n" + + "\x1dpersistent_keepalive_interval\x18\x05 \x01(\x03R\x1bpersistentKeepaliveInterval\"\xa9\x01\n" + + "\fDeviceConfig\x12\x1f\n" + + "\vprivate_key\x18\x01 \x01(\fR\n" + + "privateKey\x12\x1f\n" + + "\vlisten_port\x18\x03 \x01(\rR\n" + + "listenPort\x12E\n" + + "\x05peers\x18\x04 \x03(\v2/.v2ray.core.proxy.wireguard.wgcommon.PeerConfigR\x05peers\x12\x10\n" + + "\x03mtu\x18\x05 \x01(\rR\x03mtuB\x8a\x01\n" + + "'com.v2ray.core.proxy.wireguard.wgcommonP\x01Z7github.com/v2fly/v2ray-core/v5/proxy/wireguard/wgcommon\xaa\x02#V2Ray.Core.Proxy.Wireguard.Wgcommonb\x06proto3" + +var ( + file_proxy_wireguard_wgcommon_config_proto_rawDescOnce sync.Once + file_proxy_wireguard_wgcommon_config_proto_rawDescData []byte +) + +func file_proxy_wireguard_wgcommon_config_proto_rawDescGZIP() []byte { + file_proxy_wireguard_wgcommon_config_proto_rawDescOnce.Do(func() { + file_proxy_wireguard_wgcommon_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_proxy_wireguard_wgcommon_config_proto_rawDesc), len(file_proxy_wireguard_wgcommon_config_proto_rawDesc))) + }) + return file_proxy_wireguard_wgcommon_config_proto_rawDescData +} + +var file_proxy_wireguard_wgcommon_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2) +var file_proxy_wireguard_wgcommon_config_proto_goTypes = []any{ + (*PeerConfig)(nil), // 0: v2ray.core.proxy.wireguard.wgcommon.PeerConfig + (*DeviceConfig)(nil), // 1: v2ray.core.proxy.wireguard.wgcommon.DeviceConfig +} +var file_proxy_wireguard_wgcommon_config_proto_depIdxs = []int32{ + 0, // 0: v2ray.core.proxy.wireguard.wgcommon.DeviceConfig.peers:type_name -> v2ray.core.proxy.wireguard.wgcommon.PeerConfig + 1, // [1:1] is the sub-list for method output_type + 1, // [1:1] is the sub-list for method input_type + 1, // [1:1] is the sub-list for extension type_name + 1, // [1:1] is the sub-list for extension extendee + 0, // [0:1] is the sub-list for field type_name +} + +func init() { file_proxy_wireguard_wgcommon_config_proto_init() } +func file_proxy_wireguard_wgcommon_config_proto_init() { + if File_proxy_wireguard_wgcommon_config_proto != nil { + return + } + type x struct{} + out := protoimpl.TypeBuilder{ + File: protoimpl.DescBuilder{ + GoPackagePath: reflect.TypeOf(x{}).PkgPath(), + RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_wireguard_wgcommon_config_proto_rawDesc), len(file_proxy_wireguard_wgcommon_config_proto_rawDesc)), + NumEnums: 0, + NumMessages: 2, + NumExtensions: 0, + NumServices: 0, + }, + GoTypes: file_proxy_wireguard_wgcommon_config_proto_goTypes, + DependencyIndexes: file_proxy_wireguard_wgcommon_config_proto_depIdxs, + MessageInfos: file_proxy_wireguard_wgcommon_config_proto_msgTypes, + }.Build() + File_proxy_wireguard_wgcommon_config_proto = out.File + file_proxy_wireguard_wgcommon_config_proto_goTypes = nil + file_proxy_wireguard_wgcommon_config_proto_depIdxs = nil +} diff --git a/proxy/wireguard/wgcommon/config.proto b/proxy/wireguard/wgcommon/config.proto new file mode 100644 index 000000000..7bd219be4 --- /dev/null +++ b/proxy/wireguard/wgcommon/config.proto @@ -0,0 +1,24 @@ +syntax = "proto3"; + +package v2ray.core.proxy.wireguard.wgcommon; +option csharp_namespace = "V2Ray.Core.Proxy.Wireguard.Wgcommon"; +option go_package = "github.com/v2fly/v2ray-core/v5/proxy/wireguard/wgcommon"; +option java_package = "com.v2ray.core.proxy.wireguard.wgcommon"; +option java_multiple_files = true; + + +message PeerConfig { + bytes public_key = 1; + bytes preshared_key = 2; + repeated string allowed_ips = 3; + string endpoint = 4; + int64 persistent_keepalive_interval = 5; +} + + +message DeviceConfig { + bytes private_key = 1; + uint32 listen_port = 3; + repeated PeerConfig peers = 4; + uint32 mtu = 5; +} \ No newline at end of file diff --git a/proxy/wireguard/wgcommon/errors.generated.go b/proxy/wireguard/wgcommon/errors.generated.go new file mode 100644 index 000000000..375a9d7e0 --- /dev/null +++ b/proxy/wireguard/wgcommon/errors.generated.go @@ -0,0 +1,9 @@ +package wgcommon + +import "github.com/v2fly/v2ray-core/v5/common/errors" + +type errPathObjHolder struct{} + +func newError(values ...interface{}) *errors.Error { + return errors.New(values...).WithPathObj(errPathObjHolder{}) +} diff --git a/proxy/wireguard/wgcommon/setup.go b/proxy/wireguard/wgcommon/setup.go new file mode 100644 index 000000000..9897b1388 --- /dev/null +++ b/proxy/wireguard/wgcommon/setup.go @@ -0,0 +1,121 @@ +package wgcommon + +import ( + "errors" + "fmt" + "strings" + + "golang.zx2c4.com/wireguard/device" +) + +func (w *WrappedWireguardDevice) InitDevice() error { + if w == nil || w.config == nil { + return errors.New("wireguard: missing config") + } + if w.device != nil { + return errors.New("wireguard: device already initialized") + } + + // Create a tun device adaptor from the packetswitch network layer device. + // Use a reasonable default MTU and batch sizes. These can be tuned later. + tunDev, err := NewNetworkLayerDeviceToWireguardTunDeviceAdaptor(int(w.config.Mtu), w.tunnel, 1, 1024) + if err != nil { + return err + } + + // Create wireguard bind adapter from provided PacketConn + bind := NewNetPacketConnToWg(w.conn) + + // Create the wireguard device with our logger adapter. + dev := device.NewDevice(tunDev, bind, NewDeviceLoggerAdapter()) + if dev == nil { + return errors.New("wireguard: failed to initialize device") + } + w.device = dev + return nil +} + +func (w *WrappedWireguardDevice) SetupDeviceWithoutPeers() error { + if w == nil || w.config == nil { + return errors.New("wireguard: missing config") + } + if w.device == nil { + return errors.New("wireguard: device not initialized") + } + + var sb strings.Builder + if len(w.config.PrivateKey) > 0 { + sb.WriteString(fmt.Sprintf("private_key=%x\n", w.config.PrivateKey)) + } + if w.config.ListenPort != 0 { + sb.WriteString(fmt.Sprintf("listen_port=%d\n", w.config.ListenPort)) + } + + // Terminate operation with a blank line. + sb.WriteString("\n") + + return w.device.IpcSet(sb.String()) +} + +func (w *WrappedWireguardDevice) AddOrReplacePeers(peers []*PeerConfig) error { + if w == nil || w.config == nil { + return errors.New("wireguard: missing config") + } + if w.device == nil { + return errors.New("wireguard: device not initialized") + } + + var sb strings.Builder + // Replace existing peers with the provided list + sb.WriteString("replace_peers=true\n") + + for _, p := range peers { + if p == nil || len(p.PublicKey) == 0 { + // skip empty entries + continue + } + // start peer block + sb.WriteString(fmt.Sprintf("public_key=%x\n", p.PublicKey)) + if len(p.PresharedKey) > 0 { + sb.WriteString(fmt.Sprintf("preshared_key=%x\n", p.PresharedKey)) + } + if p.Endpoint != "" { + sb.WriteString(fmt.Sprintf("endpoint=%s\n", p.Endpoint)) + } + if p.PersistentKeepaliveInterval != 0 { + sb.WriteString(fmt.Sprintf("persistent_keepalive_interval=%d\n", p.PersistentKeepaliveInterval)) + } + // replace allowed IPs for this peer + sb.WriteString("replace_allowed_ips=true\n") + for _, aip := range p.AllowedIps { + if aip == "" { + continue + } + sb.WriteString(fmt.Sprintf("allowed_ip=%s\n", aip)) + } + } + + // terminate + sb.WriteString("\n") + + return w.device.IpcSet(sb.String()) +} + +func (w *WrappedWireguardDevice) RemovePeer(publicKey []byte) error { + if w == nil { + return errors.New("wireguard: nil receiver") + } + if w.device == nil { + return errors.New("wireguard: device not initialized") + } + if len(publicKey) == 0 { + return errors.New("wireguard: empty public key") + } + + var sb strings.Builder + sb.WriteString(fmt.Sprintf("public_key=%x\n", publicKey)) + sb.WriteString("remove=true\n") + sb.WriteString("\n") + + return w.device.IpcSet(sb.String()) +} diff --git a/proxy/wireguard/wgcommon/wgConnAdaptor.go b/proxy/wireguard/wgcommon/wgConnAdaptor.go new file mode 100644 index 000000000..b47f1ea63 --- /dev/null +++ b/proxy/wireguard/wgcommon/wgConnAdaptor.go @@ -0,0 +1,181 @@ +package wgcommon + +import ( + gonet "net" + "net/netip" + "sync" + "time" + + "github.com/v2fly/v2ray-core/v5/common/net" + "golang.zx2c4.com/wireguard/conn" +) + +// netPacketConnToWg is machine generated +type netPacketConnToWg struct { + mu sync.Mutex + conn net.PacketConn + actualPort uint16 + closed bool // tracks whether the bind is logically closed (not the conn) +} + +// NewNetPacketConnToWg constructs a wireguard conn.Bind adapter from a +// common/net.PacketConn. It returns a Bind implementation that delegates +// reads/writes to the provided PacketConn. +// +// Important: the Bind does NOT own the PacketConn lifecycle. WireGuard calls +// Close() + Open() internally during BindUpdate(); Close() here only marks +// the bind as logically closed without closing the underlying conn, so that +// Open() can re-use it. +func NewNetPacketConnToWg(c net.PacketConn) conn.Bind { + if c == nil { + return &netPacketConnToWg{} + } + n := &netPacketConnToWg{conn: c, closed: true} + if la := c.LocalAddr(); la != nil { + if ua, ok := la.(*gonet.UDPAddr); ok { + n.actualPort = uint16(ua.Port) + } + } + return n +} + +// wgEndpoint is a minimal implementation of conn.Endpoint backed by netip.AddrPort. +type wgEndpoint struct { + ap netip.AddrPort + hasSrc bool + srcIP netip.Addr +} + +func (e *wgEndpoint) ClearSrc() { + e.hasSrc = false +} + +func (e *wgEndpoint) SrcToString() string { + if !e.hasSrc { + return "" + } + // return just IP (no port) if src port is unknown + return e.srcIP.String() +} + +func (e *wgEndpoint) DstToString() string { + return e.ap.String() +} + +func (e *wgEndpoint) DstToBytes() []byte { + b, _ := e.ap.MarshalBinary() + return b +} + +func (e *wgEndpoint) DstIP() netip.Addr { + return e.ap.Addr() +} + +func (e *wgEndpoint) SrcIP() netip.Addr { + if e.hasSrc { + return e.srcIP + } + return netip.Addr{} +} + +func (n *netPacketConnToWg) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err error) { + if n.conn == nil { + return nil, 0, nil + } + n.mu.Lock() + n.closed = false + n.mu.Unlock() + + // Clear the read deadline that Close() may have set so reads can proceed. + _ = n.conn.SetReadDeadline(time.Time{}) + + // determine actualPort from LocalAddr if possible + if la := n.conn.LocalAddr(); la != nil { + if ua, ok := la.(*gonet.UDPAddr); ok { + n.actualPort = uint16(ua.Port) + } + } + actualPort = n.actualPort + + fn := func(packets [][]byte, sizes []int, eps []conn.Endpoint) (int, error) { + var i int + for i = 0; i < len(packets); i++ { + nRead, addr, err := n.conn.ReadFrom(packets[i]) + if err != nil { + if i == 0 { + return 0, err + } + return i, nil + } + sizes[i] = nRead + // build endpoint from addr + if udpAddr, ok := addr.(*gonet.UDPAddr); ok { + ip, _ := netip.AddrFromSlice(udpAddr.IP) + ap := netip.AddrPortFrom(ip, uint16(udpAddr.Port)) + eps[i] = &wgEndpoint{ap: ap} + } else { + // fallback: parse string + s := addr.String() + if ap, perr := netip.ParseAddrPort(s); perr == nil { + eps[i] = &wgEndpoint{ap: ap} + } else { + eps[i] = &wgEndpoint{} + } + } + } + return i, nil + } + return []conn.ReceiveFunc{fn}, actualPort, nil +} + +func (n *netPacketConnToWg) Close() error { + n.mu.Lock() + defer n.mu.Unlock() + n.closed = true + // Do NOT close the underlying conn here. WireGuard calls Close()+Open() + // internally during BindUpdate(). The actual PacketConn lifecycle is + // managed externally by the session that created it. + // + // Set a past read deadline to unblock any pending ReadFrom calls in + // receive goroutines so that WireGuard's stopping.Wait() can complete. + if n.conn != nil { + _ = n.conn.SetReadDeadline(time.Unix(1, 0)) + } + return nil +} + +func (n *netPacketConnToWg) SetMark(mark uint32) error { + // best-effort: underlying PacketConn may not support setting fwmark; ignore. + return nil +} + +func (n *netPacketConnToWg) Send(bufs [][]byte, ep conn.Endpoint) error { + if n.conn == nil { + return nil + } + // Use DstToString to obtain "ip:port" and resolve to UDPAddr + addrStr := ep.DstToString() + udpAddr, err := gonet.ResolveUDPAddr("udp", addrStr) + if err != nil { + return err + } + for _, b := range bufs { + if _, werr := n.conn.WriteTo(b, udpAddr); werr != nil { + return werr + } + } + return nil +} + +func (n *netPacketConnToWg) ParseEndpoint(s string) (conn.Endpoint, error) { + ap, err := netip.ParseAddrPort(s) + if err != nil { + return nil, err + } + return &wgEndpoint{ap: ap}, nil +} + +func (n *netPacketConnToWg) BatchSize() int { + // underlying common/net.PacketConn may not support batch; report 1. + return 1 +} diff --git a/proxy/wireguard/wgcommon/wgConnAdaptor_test.go b/proxy/wireguard/wgcommon/wgConnAdaptor_test.go new file mode 100644 index 000000000..c9ec8644c --- /dev/null +++ b/proxy/wireguard/wgcommon/wgConnAdaptor_test.go @@ -0,0 +1,250 @@ +package wgcommon + +import ( + "net" + "testing" + "time" + + "golang.zx2c4.com/wireguard/conn" +) + +func TestNetPacketConnToWg_OpenReceive_Send(t *testing.T) { + // setup a UDP listener (server) + svAddr, err := net.ResolveUDPAddr("udp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + svConn, err := net.ListenUDP("udp", svAddr) + if err != nil { + t.Fatal(err) + } + defer func() { _ = svConn.Close() }() + + // client + clConn, err := net.DialUDP("udp", nil, svConn.LocalAddr().(*net.UDPAddr)) + if err != nil { + t.Fatal(err) + } + defer func() { _ = clConn.Close() }() + + // Wrap server conn as Bind + bind := NewNetPacketConnToWg(svConn) + fns, port, err := bind.Open(0) + if err != nil { + t.Fatal(err) + } + if port == 0 { + // LocalAddr should have set actualPort, otherwise use svConn + if la := svConn.LocalAddr(); la != nil { + if ua, ok := la.(*net.UDPAddr); ok { + port = uint16(ua.Port) + } + } + } + if len(fns) == 0 { + t.Fatal("no receive functions returned") + } + + recvFn := fns[0] + + // run receiver in a goroutine + recvBuf := make([]byte, 1500) + sizes := make([]int, 1) + eps := make([]conn.Endpoint, 1) + ch := make(chan error, 1) + go func() { + _, err := recvFn([][]byte{recvBuf}, sizes, eps) + ch <- err + }() + + // send a message from client to server + msg := []byte("hello-wg") + if _, err := clConn.Write(msg); err != nil { + t.Fatal(err) + } + + // wait for receive + select { + case err := <-ch: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("timeout waiting for receive") + } + + // verify sizes and endpoint + if sizes[0] != len(msg) { + t.Fatalf("unexpected size: got %d want %d", sizes[0], len(msg)) + } + if eps[0] == nil { + t.Fatal("nil endpoint returned") + } + // endpoint DstToString should be the client's address + if eps[0].DstToString() != clConn.LocalAddr().String() { + t.Fatalf("unexpected endpoint dst: %s", eps[0].DstToString()) + } + + // Test Send: use a separate client conn to receive + rcv2, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = rcv2.Close() }() + + // create a bind adapter from a separate unconnected "sender" socket + senderConn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = senderConn.Close() }() + senderBind := NewNetPacketConnToWg(senderConn) + // parse endpoint for rcv2 + ep, err := senderBind.ParseEndpoint(rcv2.LocalAddr().String()) + if err != nil { + t.Fatal(err) + } + + // send from senderBind to rcv2 via Send + p := [][]byte{[]byte("ping")} + if err := senderBind.Send(p, ep); err != nil { + t.Fatal(err) + } + + // read on rcv2 + buf := make([]byte, 1500) + if err := rcv2.SetReadDeadline(time.Now().Add(time.Second)); err != nil { + t.Fatal(err) + } + n, _, err := rcv2.ReadFromUDP(buf) + if err != nil { + t.Fatal(err) + } + if n != len(p[0]) { + t.Fatalf("unexpected recv len: %d", n) + } + _ = buf +} + +func TestParseEndpointAndBatchSizeAndNilConstructor(t *testing.T) { + // Parse valid IPv4 endpoint + ep, err := NewNetPacketConnToWg(nil).ParseEndpoint("127.0.0.1:12345") + if err != nil { + t.Fatalf("ParseEndpoint failed: %v", err) + } + if ep == nil { + t.Fatal("expected endpoint, got nil") + } + if ep.DstToString() != "127.0.0.1:12345" { + t.Fatalf("unexpected DstToString: %s", ep.DstToString()) + } + + // Parse IPv6 endpoint + ipv6ep, err := NewNetPacketConnToWg(nil).ParseEndpoint("[::1]:54321") + if err != nil { + t.Fatalf("ParseEndpoint v6 failed: %v", err) + } + if ipv6ep == nil { + t.Fatal("expected ipv6 endpoint, got nil") + } + + // BatchSize and Close/SetMark for nil-constructed adapter + bind := NewNetPacketConnToWg(nil) + if bind.BatchSize() != 1 { + t.Fatalf("unexpected batch size: %d", bind.BatchSize()) + } + // Close and SetMark should be no-ops and not panic + if err := bind.Close(); err != nil { + t.Fatalf("Close on nil adapter returned error: %v", err) + } + if err := bind.SetMark(123); err != nil { + t.Fatalf("SetMark on nil adapter returned error: %v", err) + } +} + +func TestSendWithConnectedSocketProducesError(t *testing.T) { + // create a server to target + target, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = target.Close() }() + + // create a connected client socket (DialUDP) + cl, err := net.DialUDP("udp", nil, target.LocalAddr().(*net.UDPAddr)) + if err != nil { + t.Fatal(err) + } + defer func() { _ = cl.Close() }() + + bind := NewNetPacketConnToWg(cl) + ep, err := bind.ParseEndpoint("127.0.0.1:1") + if err != nil { + t.Fatal(err) + } + // Attempting to Send with a connected socket uses WriteTo and is expected + // to return an error about using WriteTo on a pre-connected connection. + p := [][]byte{[]byte("x")} + err = bind.Send(p, ep) + if err == nil { + t.Fatalf("expected error when calling Send on adapter wrapping connected socket, got nil") + } +} + +func TestReceiveMultiplePackets(t *testing.T) { + // server + sv, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0}) + if err != nil { + t.Fatal(err) + } + defer func() { _ = sv.Close() }() + + // client + c, err := net.DialUDP("udp", nil, sv.LocalAddr().(*net.UDPAddr)) + if err != nil { + t.Fatal(err) + } + defer func() { _ = c.Close() }() + + bind := NewNetPacketConnToWg(sv) + fns, _, err := bind.Open(0) + if err != nil { + t.Fatal(err) + } + if len(fns) == 0 { + t.Fatal("no receive functions") + } + recv := fns[0] + + // prepare two buffers + b1 := make([]byte, 64) + b2 := make([]byte, 64) + sizes := make([]int, 2) + eps := make([]conn.Endpoint, 2) + ch := make(chan error, 1) + go func() { + _, err := recv([][]byte{b1, b2}, sizes, eps) + ch <- err + }() + + // send two packets quickly + if _, err := c.Write([]byte("one")); err != nil { + t.Fatal(err) + } + if _, err := c.Write([]byte("two")); err != nil { + t.Fatal(err) + } + + select { + case err := <-ch: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("timeout waiting for batched receive") + } + + if sizes[0] == 0 && sizes[1] == 0 { + t.Fatalf("expected at least one packet size to be non-zero") + } +} diff --git a/proxy/wireguard/wgcommon/wgDeviceAdaptor.go b/proxy/wireguard/wgcommon/wgDeviceAdaptor.go new file mode 100644 index 000000000..fba318b28 --- /dev/null +++ b/proxy/wireguard/wgcommon/wgDeviceAdaptor.go @@ -0,0 +1,201 @@ +package wgcommon + +import ( + "errors" + "os" + "sync" + + "github.com/v2fly/v2ray-core/v5/common/packetswitch" + "golang.zx2c4.com/wireguard/tun" +) + +func NewNetworkLayerDeviceToWireguardTunDeviceAdaptor(mtu int, networkLayerSwitch packetswitch.NetworkLayerDevice, batchSize int, inboundChannelSize int) (*NetworkLayerDeviceToWireguardTunDeviceAdaptor, error) { + if batchSize <= 0 { + batchSize = 1 + } + if inboundChannelSize <= 0 { + inboundChannelSize = 1024 + } + n := &NetworkLayerDeviceToWireguardTunDeviceAdaptor{ + mtu: mtu, + networkLayerSwitch: networkLayerSwitch, + in: make(chan []byte, inboundChannelSize), + events: make(chan tun.Event, 4), + batchSize: batchSize, + } + // Attach writer to the network layer switch so incoming packets are delivered to this adaptor. + if networkLayerSwitch != nil { + if err := networkLayerSwitch.OnAttach(&networkLayerWriter{parent: n}); err != nil { + return nil, err + } + // If the underlying device exposes real link events, forward them. + if src, ok := networkLayerSwitch.(interface{ Events() <-chan tun.Event }); ok { + go func() { + for ev := range src.Events() { + // Acquire lock to synchronize with Close (which closes the events channel). + n.mu.Lock() + closed := n.closed + if !closed { + // best-effort: do not block if events channel is full + select { + case n.events <- ev: + default: + } + } + n.mu.Unlock() + if closed { + return + } + } + }() + } + } + return n, nil +} + +// NetworkLayerDeviceToWireguardTunDeviceAdaptor is primarily machine generated. +type NetworkLayerDeviceToWireguardTunDeviceAdaptor struct { + mtu int + networkLayerSwitch packetswitch.NetworkLayerDevice + + mu sync.RWMutex + closed bool + in chan []byte + events chan tun.Event + batchSize int +} + +// networkLayerWriter adapts packetswitch writes into the adaptor's incoming channel. +type networkLayerWriter struct { + parent *NetworkLayerDeviceToWireguardTunDeviceAdaptor +} + +func (w *networkLayerWriter) Write(packet []byte) (int, error) { + p := make([]byte, len(packet)) + copy(p, packet) + w.parent.mu.RLock() + closed := w.parent.closed + w.parent.mu.RUnlock() + if closed { + return 0, errors.New("device closed") + } + select { + case w.parent.in <- p: + return len(packet), nil + default: + // Channel full, drop packet. + return 0, errors.New("no buffer space") + } +} + +func (n *NetworkLayerDeviceToWireguardTunDeviceAdaptor) File() *os.File { + // No underlying OS file for this adaptor. + return nil +} + +func (n *NetworkLayerDeviceToWireguardTunDeviceAdaptor) Read(bufs [][]byte, sizes []int, offset int) (ret int, err error) { + // Read up to BatchSize packets or until bufs exhausted. + // NOTE: 'offset' is a byte offset within each buffer (to leave room for transport headers), + // NOT an index into the bufs slice. WireGuard expects packet data to be written starting at + // bufs[i][offset:]. + maxCount := n.BatchSize() + if maxCount <= 0 { + maxCount = 1 + } + for i := 0; i < maxCount && i < len(bufs); i++ { + var b []byte + var ok bool + if ret == 0 { + // first read: block until a packet arrives or channel closes + b, ok = <-n.in + } else { + // subsequent reads: do not block — if no packet available, return what we've got + select { + case b, ok = <-n.in: + // got one + default: + return ret, nil + } + } + if !ok { + // channel closed + if ret == 0 { + return 0, os.ErrClosed + } + return ret, nil + } + to := bufs[i] + if to == nil { + // packet consumed but no destination buffer provided, skip copying + ret++ + continue + } + copied := copy(to[offset:], b) + if sizes != nil && i < len(sizes) { + sizes[i] = copied + } + ret++ + } + return ret, nil +} + +func (n *NetworkLayerDeviceToWireguardTunDeviceAdaptor) Write(bufs [][]byte, offset int) (int, error) { + written := 0 + if n.networkLayerSwitch == nil { + return 0, errors.New("no network layer writer attached") + } + for i := 0; i < len(bufs); i++ { + b := bufs[i] + if b == nil || len(b) <= offset { + continue + } + // The offset is a byte offset within each buffer where the actual + // packet payload starts (after transport headers). Extract only the + // payload portion. Copy because caller may reuse buffer. + payload := b[offset:] + cp := make([]byte, len(payload)) + copy(cp, payload) + _, err := n.networkLayerSwitch.Write(cp) + if err != nil { + return written, err + } + written++ + } + return written, nil +} + +func (n *NetworkLayerDeviceToWireguardTunDeviceAdaptor) MTU() (int, error) { + return n.mtu, nil +} + +func (n *NetworkLayerDeviceToWireguardTunDeviceAdaptor) Name() (string, error) { + // No specific name available for this virtual adaptor. + return "", nil +} + +func (n *NetworkLayerDeviceToWireguardTunDeviceAdaptor) Events() <-chan tun.Event { + return n.events +} + +func (n *NetworkLayerDeviceToWireguardTunDeviceAdaptor) Close() error { + n.mu.Lock() + if n.closed { + n.mu.Unlock() + return nil + } + n.closed = true + close(n.in) + // Close events channel to signal no more events. + close(n.events) + dev := n.networkLayerSwitch + n.networkLayerSwitch = nil + n.mu.Unlock() + if dev != nil { + _ = dev.Close() + } + return nil +} + +func (n *NetworkLayerDeviceToWireguardTunDeviceAdaptor) BatchSize() int { + return n.batchSize +} diff --git a/proxy/wireguard/wgcommon/wgDeviceAdaptor_test.go b/proxy/wireguard/wgcommon/wgDeviceAdaptor_test.go new file mode 100644 index 000000000..7cf922a33 --- /dev/null +++ b/proxy/wireguard/wgcommon/wgDeviceAdaptor_test.go @@ -0,0 +1,276 @@ +package wgcommon + +import ( + "reflect" + "sync" + "testing" + "time" + + "github.com/v2fly/v2ray-core/v5/common/packetswitch" + "golang.zx2c4.com/wireguard/tun" +) + +// fakeNetDevice implements packetswitch.NetworkLayerDevice and optionally exposes Events(). +type fakeNetDevice struct { + mu sync.Mutex + writer packetswitch.NetworkLayerPacketWriter + writes [][]byte + closed bool + events chan tun.Event +} + +func (f *fakeNetDevice) OnAttach(w packetswitch.NetworkLayerPacketWriter) error { + f.mu.Lock() + defer f.mu.Unlock() + f.writer = w + return nil +} + +func (f *fakeNetDevice) Write(packet []byte) (int, error) { + f.mu.Lock() + defer f.mu.Unlock() + if f.closed { + return 0, errClosed + } + cp := make([]byte, len(packet)) + copy(cp, packet) + f.writes = append(f.writes, cp) + return len(packet), nil +} + +func (f *fakeNetDevice) Close() error { + f.mu.Lock() + f.closed = true + f.mu.Unlock() + return nil +} + +func (f *fakeNetDevice) getWriter() packetswitch.NetworkLayerPacketWriter { + f.mu.Lock() + w := f.writer + f.mu.Unlock() + return w +} + +func (f *fakeNetDevice) lastWrite() []byte { + f.mu.Lock() + defer f.mu.Unlock() + if len(f.writes) == 0 { + return nil + } + return f.writes[len(f.writes)-1] +} + +// Provide Events() so adaptor can forward events when present. +func (f *fakeNetDevice) Events() <-chan tun.Event { + return f.events +} + +var errClosed = &fakeError{"closed"} + +type fakeError struct{ s string } + +func (e *fakeError) Error() string { return e.s } + +func TestNewAdaptor_ReadWrite_Basic(t *testing.T) { + fd := &fakeNetDevice{events: make(chan tun.Event, 4)} + // batchSize 2, inboundChannelSize 4 + a, err := NewNetworkLayerDeviceToWireguardTunDeviceAdaptor(1500, fd, 2, 4) + if err != nil { + t.Fatalf("constructor failed: %v", err) + } + + w := fd.getWriter() + if w == nil { + t.Fatal("expected writer to be attached to fake device") + } + + p1 := []byte{0x01, 0x02, 0x03} + if _, err := w.Write(p1); err != nil { + t.Fatalf("writer.Write failed: %v", err) + } + + bufs := make([][]byte, 2) + bufs[0] = make([]byte, 64) + bufs[1] = make([]byte, 64) + sizes := make([]int, 2) + + ret, err := a.Read(bufs, sizes, 0) + if err != nil { + t.Fatalf("Read returned error: %v", err) + } + if ret != 1 { + t.Fatalf("expected 1 packet read, got %d", ret) + } + if sizes[0] != len(p1) { + t.Fatalf("expected sizes[0]=%d, got %d", len(p1), sizes[0]) + } + if !reflect.DeepEqual(bufs[0][:sizes[0]], p1) { + t.Fatalf("payload mismatch: got %v want %v", bufs[0][:sizes[0]], p1) + } + + // Now write two packets and read both + p2 := []byte{0x0a, 0x0b} + p3 := []byte{0x0c} + if _, err := w.Write(p2); err != nil { + t.Fatalf("writer.Write p2 failed: %v", err) + } + if _, err := w.Write(p3); err != nil { + t.Fatalf("writer.Write p3 failed: %v", err) + } + + bufs2 := make([][]byte, 2) + bufs2[0] = make([]byte, 16) + bufs2[1] = make([]byte, 16) + sizes2 := make([]int, 2) + + ret2, err := a.Read(bufs2, sizes2, 0) + if err != nil { + t.Fatalf("Read returned error: %v", err) + } + if ret2 != 2 { + t.Fatalf("expected 2 packets read, got %d", ret2) + } + if sizes2[0] != len(p2) || sizes2[1] != len(p3) { + t.Fatalf("unexpected sizes: %v", sizes2) + } +} + +func TestNewAdaptor_WriteToNetwork(t *testing.T) { + fd := &fakeNetDevice{events: make(chan tun.Event, 4)} + a, err := NewNetworkLayerDeviceToWireguardTunDeviceAdaptor(1500, fd, 1, 4) + if err != nil { + t.Fatalf("constructor failed: %v", err) + } + + payload := []byte{0xaa, 0xbb, 0xcc} + bufs := make([][]byte, 1) + bufs[0] = payload + + written, err := a.Write(bufs, 0) + if err != nil { + t.Fatalf("Write returned error: %v", err) + } + if written != 1 { + t.Fatalf("expected 1 written, got %d", written) + } + lw := fd.lastWrite() + if !reflect.DeepEqual(lw, payload) { + t.Fatalf("device write mismatch: got %v want %v", lw, payload) + } +} + +func TestNewAdaptor_InboundBufferFullDrops(t *testing.T) { + fd := &fakeNetDevice{events: make(chan tun.Event, 4)} + // inboundChannelSize 1 so second write should fail + a, err := NewNetworkLayerDeviceToWireguardTunDeviceAdaptor(1500, fd, 2, 1) + if err != nil { + t.Fatalf("constructor failed: %v", err) + } + w := fd.getWriter() + if w == nil { + t.Fatal("expected writer to be attached to fake device") + } + + p1 := []byte{0x01} + p2 := []byte{0x02} + if _, err := w.Write(p1); err != nil { + t.Fatalf("first write failed: %v", err) + } + // second write should return error due to buffer full + if _, err := w.Write(p2); err == nil { + t.Fatalf("expected second write to fail due to full buffer") + } + + bufs := make([][]byte, 1) + bufs[0] = make([]byte, 8) + sizes := make([]int, 1) + ret, err := a.Read(bufs, sizes, 0) + if err != nil { + t.Fatalf("Read returned error: %v", err) + } + if ret != 1 { + t.Fatalf("expected 1 packet read, got %d", ret) + } +} + +func TestNewAdaptor_ReadWithNonZeroOffset(t *testing.T) { + // This test reproduces the 100% CPU bug where offset was misinterpreted + // as a buffer index rather than a byte offset within each buffer. + fd := &fakeNetDevice{events: make(chan tun.Event, 4)} + a, err := NewNetworkLayerDeviceToWireguardTunDeviceAdaptor(1500, fd, 1, 4) + if err != nil { + t.Fatalf("constructor failed: %v", err) + } + w := fd.getWriter() + if w == nil { + t.Fatal("expected writer to be attached") + } + + p1 := []byte{0xDE, 0xAD, 0xBE, 0xEF} + if _, err := w.Write(p1); err != nil { + t.Fatalf("writer.Write failed: %v", err) + } + + // Use offset=16 to mimic WireGuard's MessageTransportHeaderSize + const offset = 16 + bufs := make([][]byte, 1) + bufs[0] = make([]byte, 64) + sizes := make([]int, 1) + + ret, err := a.Read(bufs, sizes, offset) + if err != nil { + t.Fatalf("Read returned error: %v", err) + } + if ret != 1 { + t.Fatalf("expected 1 packet read, got %d", ret) + } + if sizes[0] != len(p1) { + t.Fatalf("expected sizes[0]=%d, got %d", len(p1), sizes[0]) + } + // Data should be at bufs[0][offset:offset+sizes[0]], not bufs[0][0:sizes[0]] + got := bufs[0][offset : offset+sizes[0]] + if !reflect.DeepEqual(got, p1) { + t.Fatalf("payload mismatch: got %v want %v", got, p1) + } + // The header area before offset should be untouched (all zeros) + for i := 0; i < offset; i++ { + if bufs[0][i] != 0 { + t.Fatalf("byte at position %d was modified: %x", i, bufs[0][i]) + } + } +} + +func TestNewAdaptor_EventForwardingAndClose(t *testing.T) { + fd := &fakeNetDevice{events: make(chan tun.Event, 4)} + a, err := NewNetworkLayerDeviceToWireguardTunDeviceAdaptor(1500, fd, 1, 4) + if err != nil { + t.Fatalf("constructor failed: %v", err) + } + + // send event from underlying device + fd.events <- tun.EventUp + + select { + case ev := <-a.Events(): + if ev != tun.EventUp { + t.Fatalf("expected EventUp, got %v", ev) + } + case <-time.After(time.Second): + t.Fatal("timeout waiting for forwarded event") + } + + // Close adaptor and ensure events channel is closed + if err := a.Close(); err != nil { + t.Fatalf("Close returned error: %v", err) + } + // reading from closed events channel should return immediately with ok==false + select { + case _, ok := <-a.Events(): + if ok { + t.Fatal("expected events channel to be closed") + } + case <-time.After(time.Second): + t.Fatal("timeout waiting for events channel close") + } +} diff --git a/proxy/wireguard/wgcommon/wgLogAdaptor.go b/proxy/wireguard/wgcommon/wgLogAdaptor.go new file mode 100644 index 000000000..35388d41e --- /dev/null +++ b/proxy/wireguard/wgcommon/wgLogAdaptor.go @@ -0,0 +1,27 @@ +package wgcommon + +import ( + "fmt" + + "github.com/v2fly/v2ray-core/v5/common/errors" + "golang.zx2c4.com/wireguard/device" +) + +// NewDeviceLoggerAdapter returns a wireguard device.Logger that forwards +// verbose and error logs into the project's error logger using errors.New(...). +// Verbosef logs are recorded as Debug, Errorf logs are recorded as Error. +// machine generated +func NewDeviceLoggerAdapter() *device.Logger { + l := &device.Logger{} + l.Verbosef = func(format string, args ...any) { + msg := fmt.Sprintf(format, args...) + err := errors.New(msg) + err.AtDebug().WriteToLog() + } + l.Errorf = func(format string, args ...any) { + msg := fmt.Sprintf(format, args...) + err := errors.New(msg) + err.AtError().WriteToLog() + } + return l +} diff --git a/proxy/wireguard/wgcommon/wgcommon.go b/proxy/wireguard/wgcommon/wgcommon.go new file mode 100644 index 000000000..b5f792ea9 --- /dev/null +++ b/proxy/wireguard/wgcommon/wgcommon.go @@ -0,0 +1,3 @@ +package wgcommon + +//go:generate go run github.com/v2fly/v2ray-core/v5/common/errors/errorgen diff --git a/proxy/wireguard/wgcommon/wgdevice.go b/proxy/wireguard/wgcommon/wgdevice.go new file mode 100644 index 000000000..744fbe5bd --- /dev/null +++ b/proxy/wireguard/wgcommon/wgdevice.go @@ -0,0 +1,71 @@ +package wgcommon + +import ( + "context" + + "github.com/v2fly/v2ray-core/v5/common/net" + "github.com/v2fly/v2ray-core/v5/common/packetswitch" + "golang.zx2c4.com/wireguard/device" +) + +func NewWrappedWireguardDevice(ctx context.Context, config *DeviceConfig) (*WrappedWireguardDevice, error) { + return &WrappedWireguardDevice{ + config: config, + ctx: ctx, + }, nil +} + +type WrappedWireguardDevice struct { + config *DeviceConfig + ctx context.Context + device *device.Device + + tunnel packetswitch.NetworkLayerDevice + conn net.PacketConn +} + +func (w *WrappedWireguardDevice) Up() error { + if w.device != nil { + return w.device.Up() + } + return newError("wireguard device do not exist").AtError() +} + +// SetTunnel sets the network layer tunnel device for the wrapped WireGuard device. +func (w *WrappedWireguardDevice) SetTunnel(t packetswitch.NetworkLayerDevice) { + w.tunnel = t +} + +// SetConn sets the underlying packet connection used by the wrapped WireGuard device. +func (w *WrappedWireguardDevice) SetConn(c net.PacketConn) { + w.conn = c +} + +func (w *WrappedWireguardDevice) Close() error { + if w == nil { + return nil + } + // Bring device down if initialized + if w.device != nil { + _ = w.device.Down() + w.device = nil + } + // Close tunnel if present + if w.tunnel != nil { + _ = w.tunnel.Close() + w.tunnel = nil + } + // Close underlying packet conn if present + if w.conn != nil { + _ = w.conn.Close() + w.conn = nil + } + return nil +} + +func (w *WrappedWireguardDevice) Debug() (string, error) { + if w.device != nil { + return w.device.IpcGet() + } + return "", nil +} diff --git a/transport/internet/tlsmirror/mirrorenrollment/clicommand/enrollmentlink_cli.go b/transport/internet/tlsmirror/mirrorenrollment/clicommand/enrollmentlink_cli.go index 496dca220..208dfdc9f 100644 --- a/transport/internet/tlsmirror/mirrorenrollment/clicommand/enrollmentlink_cli.go +++ b/transport/internet/tlsmirror/mirrorenrollment/clicommand/enrollmentlink_cli.go @@ -19,6 +19,7 @@ var ( mode *string ) +// machine generated var cmdEnrollmentLink = &base.Command{ UsageLine: "{{.Exec}} engineering tlsmirror-enrollment-link", Flag: func() flag.FlagSet { diff --git a/transport/internet/tlsmirror/mirrorenrollment/enrollmentlink.go b/transport/internet/tlsmirror/mirrorenrollment/enrollmentlink.go index 67916ff19..94bce7042 100644 --- a/transport/internet/tlsmirror/mirrorenrollment/enrollmentlink.go +++ b/transport/internet/tlsmirror/mirrorenrollment/enrollmentlink.go @@ -19,6 +19,7 @@ const ( // data:application/vnd.v2ray.tlsmirror-enrollment;base64, // where payload is the marshaled Any message encoded with standard base64 (with padding). func LinkFromAny(a *anypb.Any) (string, error) { + // Machine generated code if a == nil { return "", newError("nil Any") } @@ -34,6 +35,7 @@ func LinkFromAny(a *anypb.Any) (string, error) { // AnyFromLink converts a string link (now primarily a data URL) back to *anypb.Any. // Accepted formats (strict): only data URLs matching the exact MIME type and base64 encoding. func AnyFromLink(link string) (*anypb.Any, error) { + // Machine generated code if link == "" { return nil, newError("empty link") }