//go:build windows package windivert import ( "encoding/binary" "errors" "runtime" "sync" "unsafe" E "github.com/sagernet/sing/common/exceptions" "golang.org/x/sys/windows" ) type Handle struct { device windows.Handle event windows.Handle closing sync.Once closeErr error addr Address recvAddrs []Address recvAddrsLen uint32 sendAddrs []Address } func Open(filter *Filter, layer Layer, priority int16, flags Flag) (*Handle, error) { err := validateOpenArgs(layer, priority, flags) if err != nil { return nil, err } if filter == nil { filter = reject() } filterBin, filterFlags, err := filter.encode() if err != nil { return nil, err } device, err := acquireDevice() if err != nil { return nil, err } event, err := windows.CreateEvent(nil, 1, 0, nil) // manual reset, unsignaled if err != nil { windows.CloseHandle(device) return nil, E.Cause(err, "windivert: create event") } h := &Handle{device: device, event: event} err = h.initialize(layer, priority, flags) if err != nil { h.Close() return nil, err } err = h.startup(filterBin, filterFlags) if err != nil { h.Close() return nil, err } return h, nil } func openDevice() (windows.Handle, error) { return windows.CreateFile( driverDevName, windows.GENERIC_READ|windows.GENERIC_WRITE, 0, nil, windows.OPEN_EXISTING, windows.FILE_ATTRIBUTE_NORMAL|windows.FILE_FLAG_OVERLAPPED, 0, ) } func validateOpenArgs(layer Layer, priority int16, flags Flag) error { if layer != LayerNetwork { return E.New("windivert: invalid layer ", uint32(layer)) } if priority < PriorityLowest || priority > PriorityHighest { return E.New("windivert: priority out of range") } const supportedFlags = FlagSniff | FlagSendOnly if flags&^supportedFlags != 0 { return E.New("windivert: unknown flag bits") } if flags&FlagSniff != 0 && flags&FlagSendOnly != 0 { return E.New("windivert: FlagSniff and FlagSendOnly are mutually exclusive") } return nil } func (h *Handle) initialize(layer Layer, priority int16, flags Flag) error { in := buildIoctlInitialize(layer, priority, flags) var outBuf [versionStructSize]byte binary.LittleEndian.PutUint64(outBuf[0:8], magicDLL) binary.LittleEndian.PutUint32(outBuf[8:12], versionMajor) binary.LittleEndian.PutUint32(outBuf[12:16], versionMinor) binary.LittleEndian.PutUint32(outBuf[16:20], uint32(unsafe.Sizeof(uintptr(0))*8)) _, err := doIoctl(h.device, ioctlInitialize, in[:], outBuf[:], h.event) if err != nil { return E.Cause(err, "windivert: initialize ioctl") } gotMagic := binary.LittleEndian.Uint64(outBuf[0:8]) if gotMagic != magicSYS { return E.New("windivert: driver magic mismatch (got ", gotMagic, ")") } gotMajor := binary.LittleEndian.Uint32(outBuf[8:12]) if gotMajor < versionMajor { gotMinor := binary.LittleEndian.Uint32(outBuf[12:16]) return E.New("windivert: driver version too old: ", gotMajor, ".", gotMinor) } return nil } func (h *Handle) startup(filterBin []byte, filterFlags uint64) error { in := buildIoctlStartup(filterFlags) _, err := doIoctl(h.device, ioctlStartup, in[:], filterBin, h.event) if err != nil { return E.Cause(err, "windivert: startup ioctl") } return nil } func (h *Handle) Recv(buf []byte) (int, Address, error) { if len(buf) == 0 { return 0, Address{}, E.New("windivert: recv: zero-length buffer") } h.addr = Address{} in := buildIoctlRecv(&h.addr) n, err := doIoctl(h.device, ioctlRecv, in[:], buf, h.event) runtime.KeepAlive(h) if err != nil { return 0, Address{}, err } return int(n), h.addr, nil } // BatchMax is WINDIVERT_BATCH_MAX: the driver caps both directions at 255 // packets per ioctl. const BatchMax = 255 const addressSize = uint32(unsafe.Sizeof(Address{})) // The driver packs packets back-to-back into buf with no padding and copies // exactly each packet's IP total length, and returns as soon as at least one // packet is available. func (h *Handle) RecvBatch(buf []byte) (int, []Address, error) { if len(buf) < MTUMax { return 0, nil, E.New("windivert: recv batch: buffer smaller than MTUMax") } if h.recvAddrs == nil { h.recvAddrs = make([]Address, BatchMax) } h.recvAddrsLen = uint32(len(h.recvAddrs)) * addressSize in := buildIoctlRecvBatch(&h.recvAddrs[0], &h.recvAddrsLen) n, err := doIoctl(h.device, ioctlRecv, in[:], buf, h.event) runtime.KeepAlive(h) if err != nil { return 0, nil, err } return int(n), h.recvAddrs[:h.recvAddrsLen/addressSize], nil } // The driver recovers packet boundaries from the IP total-length fields and // rejects the whole batch if they do not add up to len(buf). func (h *Handle) SendBatch(buf []byte, addrs []Address) (int, error) { if len(addrs) == 0 || len(addrs) > BatchMax { return 0, E.New("windivert: send batch: invalid packet count ", len(addrs)) } if len(buf) == 0 { return 0, E.New("windivert: send batch: empty buffer") } if h.sendAddrs == nil { h.sendAddrs = make([]Address, BatchMax) } copy(h.sendAddrs, addrs) in := buildIoctlSend(&h.sendAddrs[0], uint32(len(addrs))*addressSize) n, err := doIoctl(h.device, ioctlSend, in[:], buf, h.event) runtime.KeepAlive(h) if err != nil { return 0, err } return int(n), nil } // The address's Outbound flag controls whether the packet is sent toward // the wire (outbound=true) or delivered up the stack (outbound=false). func (h *Handle) Send(packet []byte, addr *Address) (int, error) { if len(packet) == 0 { return 0, E.New("windivert: send: empty packet") } if addr == nil { return 0, E.New("windivert: send: nil address") } h.addr = *addr in := buildIoctlSend(&h.addr, addressSize) n, err := doIoctl(h.device, ioctlSend, in[:], packet, h.event) runtime.KeepAlive(h) if err != nil { return 0, err } return int(n), nil } func (h *Handle) Close() error { h.closing.Do(func() { var errs []error if h.device != 0 { err := windows.CloseHandle(h.device) if err != nil { errs = append(errs, err) } h.device = 0 } if h.event != 0 { err := windows.CloseHandle(h.event) if err != nil { errs = append(errs, err) } h.event = 0 } h.closeErr = E.Errors(errs...) }) return h.closeErr } // IOCTL codes from windivert_device.h. CTL_CODE macro layout: // // (DeviceType << 16) | (Access << 14) | (Function << 2) | Method const ( fileDeviceNetwork uint32 = 0x12 accessReadWrite uint32 = 3 // FILE_READ_DATA | FILE_WRITE_DATA accessRead uint32 = 1 methodInDirect uint32 = 1 methodOutDirect uint32 = 2 ) func ctlCode(deviceType, access, function, method uint32) uint32 { return (deviceType << 16) | (access << 14) | (function << 2) | method } var ( ioctlInitialize = ctlCode(fileDeviceNetwork, accessReadWrite, 0x921, methodOutDirect) ioctlStartup = ctlCode(fileDeviceNetwork, accessReadWrite, 0x922, methodInDirect) ioctlRecv = ctlCode(fileDeviceNetwork, accessRead, 0x923, methodOutDirect) ioctlSend = ctlCode(fileDeviceNetwork, accessReadWrite, 0x924, methodInDirect) ) // Magic numbers exchanged during INITIALIZE. DLL sends magicDLL in the // version struct; driver returns magicSYS on success. const ( magicDLL uint64 = 0x4C4C447669645724 // "$WdivDLL" in LE bytes magicSYS uint64 = 0x5359537669645723 // "#WdivSYS" in LE bytes ) const ( versionMajor uint32 = 2 versionMinor uint32 = 2 ) // Size of the WINDIVERT_IOCTL union on wire (packed). const ioctlSize = 16 // Size of WINDIVERT_VERSION on wire (packed). Only the first 20 bytes // carry data; the rest is reserved zero padding. const versionStructSize = 64 // NtDeviceIoControlFile clears the event to nonsignaled before queuing each // request, so one event can be reused across calls without ResetEvent. func doIoctl(handle windows.Handle, code uint32, in []byte, out []byte, event windows.Handle) (uint32, error) { var overlapped windows.Overlapped overlapped.HEvent = event var inPtr *byte var inLen uint32 if len(in) > 0 { inPtr = &in[0] inLen = uint32(len(in)) } var outPtr *byte var outLen uint32 if len(out) > 0 { outPtr = &out[0] outLen = uint32(len(out)) } var returned uint32 err := windows.DeviceIoControl(handle, code, inPtr, inLen, outPtr, outLen, &returned, &overlapped) if err == nil { return returned, nil } if !errors.Is(err, windows.ERROR_IO_PENDING) { return 0, err } err = windows.GetOverlappedResult(handle, &overlapped, &returned, true) if err != nil { return 0, err } return returned, nil } func buildIoctlInitialize(layer Layer, priority int16, flags Flag) [ioctlSize]byte { var buf [ioctlSize]byte binary.LittleEndian.PutUint32(buf[0:4], uint32(layer)) // The driver expects priority + WINDIVERT_PRIORITY_HIGHEST (30000) so // the low range maps to non-negative integers. binary.LittleEndian.PutUint32(buf[4:8], uint32(int32(priority)+int32(PriorityHighest))) binary.LittleEndian.PutUint64(buf[8:16], uint64(flags)) return buf } func buildIoctlStartup(filterFlags uint64) [ioctlSize]byte { var buf [ioctlSize]byte binary.LittleEndian.PutUint64(buf[0:8], filterFlags) return buf } // The driver dereferences the packed pointer to write the received packet's // WINDIVERT_ADDRESS. func buildIoctlRecv(addr *Address) [ioctlSize]byte { var buf [ioctlSize]byte binary.LittleEndian.PutUint64(buf[0:8], uint64(uintptr(unsafe.Pointer(addr)))) binary.LittleEndian.PutUint64(buf[8:16], 0) return buf } // addr_len_ptr carries the Address array capacity in bytes; the driver // overwrites it with the bytes actually written (packet count × 80). func buildIoctlRecvBatch(addrs *Address, addrsLen *uint32) [ioctlSize]byte { var buf [ioctlSize]byte binary.LittleEndian.PutUint64(buf[0:8], uint64(uintptr(unsafe.Pointer(addrs)))) binary.LittleEndian.PutUint64(buf[8:16], uint64(uintptr(unsafe.Pointer(addrsLen)))) return buf } func buildIoctlSend(addrs *Address, addrsLen uint32) [ioctlSize]byte { var buf [ioctlSize]byte binary.LittleEndian.PutUint64(buf[0:8], uint64(uintptr(unsafe.Pointer(addrs)))) binary.LittleEndian.PutUint64(buf[8:16], uint64(addrsLen)) return buf }