Compare commits

..
Author SHA1 Message Date
patternihaandClaude Opus 5.5 747b153333 TUN inbound: Error out when autoSystemWfpBlockLeak or autoSystemDnsToGateway cannot apply
As asked in review, rather than run without them:

- The config is rejected, also by xray -test, for autoSystemWfpBlockLeak
  without autoSystemRoutingTable, or with "dns" but without dns, on
  Windows, and for autoSystemDnsToGateway without gateway on Linux.
- Xray does not start when the filters cannot be added, now on every
  Windows version, or when the system DNS cannot be set on Linux,
  instead of logging it and running without them.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 08:46:45 +03:30
patternihaandClaude Opus 5.5 94cd83ccb4 TUN inbound: Rename "misconfig" to "misconfigtun"
The autoSystemWfpBlockLeak value that blocks an IP version not routed to
the TUN, as asked in review; "misconfig" was never released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 07:59:39 +03:30
patternihaandClaude Opus 5.5 1c225a8041 TUN inbound: Make autoSystemWfpBlockLeak a list of the leaks to block
autoSystemWfpBlockLeak now takes ["dns", "misconfig"] instead of true:
"dns" keeps DNS inside the TUN, and "misconfig" blocks an IP version
that no route leads to the TUN, the leak of a configuration that routes
only one of them. Either can be used alone, e.g. ["dns"] to block DNS
leaks while an IP version stays out of the TUN on purpose. Unknown
values are rejected. The config field becomes a repeated string with
the same number; it was never released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 07:46:51 +03:30
patternihaandClaude Opus 5.5 3a6bdb19ba TUN inbound: autoSystemDnsToGateway falls back to an IPv6 gateway on Linux
Without an IPv4 address in gateway, the system DNS now points at the
first IPv6 gateway plus one (e.g. fc00::1/64 -> fc00::2) instead of
nothing, and the routing check before the takeover accepts IPv6
addresses for it.

The README also says what each system does without gateway: Xray
assigns no address on Linux, Windows gives the TUN link-local ones
itself, and macOS and FreeBSD use 169.254.10.1/30.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 07:18:55 +03:30
patternihaandClaude Opus 5.5 de02da553a TUN inbound: Block IPv4 too when it is not routed to the TUN on Windows
autoSystemWfpBlockLeak blocked IPv6 when the TUN had no IPv6 address or
no IPv6 route, and IPv4 never: with only IPv6 routed to the TUN, IPv4
went around it. Now each IP version is blocked when no route of it leads
to the TUN, except for loopback, DHCP, IPv6 neighbor and multicast
listener discovery, and Xray itself.

Addresses no longer count: without one of a version in gateway, Windows
gives the TUN a link-local one itself (fe80:: at once, 169.254.x.x after
some seconds), and what is routed to the TUN goes through it with that,
so a TUN with IPv6 routes but no IPv6 address had its IPv6 blocked for
nothing.

Tested on Windows 11, elevated: with only IPv6 routed, other programs'
IPv4 is denied, while loopback, a DHCP renew of Wi-Fi and Xray still
work; without gateway, IPv4 and IPv6 routed to the TUN enter it from
169.254.x.x and fe80::.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 03:07:14 +03:30
patternihaandClaude Opus 5.5 4ec4fb8aab TUN inbound: Rename to autoSystemWfpBlockLeak and autoSystemDnsToGateway
autoSystemWFP becomes autoSystemWfpBlockLeak, saying that the WFP filters
block leaks, and autoSystemDNS becomes autoSystemDnsToGateway, saying
where it points the system DNS, so that pointing the system DNS at the
gateway on Windows later would fit the same name. Their config fields
keep their numbers; neither was released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 02:31:00 +03:30
patternihaandClaude Opus 5.5 63de6135cb TUN inbound: Rename strictRoute to autoSystemWFP
It only turns on the Windows Filtering Platform filters, along with Xray
resolving its own lookups while they restrict DNS, so it is named after
what it sets up in the system, like autoSystemRoutingTable and
autoSystemDNS. The config field becomes auto_system_wfp, with the same
number; strictRoute was never released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 01:22:22 +03:30
patternihaandClaude Opus 5.5 edd916b08e TUN inbound: Keep Windows' DNS Client from sending DoH/DoT outside the TUN
With strictRoute, only port 53 was kept inside the TUN. But Windows' DNS
Client service sends the queries for an interface's DNS servers out
through that interface, whatever the routes say, and since Windows 11
and Server 2022 it may send them over HTTPS or TLS, when that is set up
for the interface (as Windows Settings does) or for the server. Those
left through the physical link.

On those versions, the DNS Client service may now only connect through
the TUN, except for its mDNS and LLMNR. The filters recognize the
service by its SID in the token of its process, as Windows Firewall's
own rules for it do. Earlier versions only query port 53, and may run
the service in one process with others, so they get no such filters.
The port 53 rule stays, for the programs that query a resolver on the
local network themselves, and for those earlier versions.

Tested on Windows 11, elevated: with DoH set on Wi-Fi per adapter, per
network profile or by global auto-upgrade, none of the DNS Client's
connections left through Wi-Fi (WFP logged the drops by the new filter),
names still resolved through the TUN, mDNS and LLMNR still went out, and
other programs were unaffected, also in a real Xray run.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 01:09:52 +03:30
patternihaandClaude Opus 5.5 eb29a4e3de TUN inbound: Make strictRoute false by default
Like sing-box's strict_route, strictRoute is now false by default, so the
Windows Filtering Platform filters are only added when it is set to true
(together with autoSystemRoutingTable). With unset meaning false,
strict_route becomes a plain bool field.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-29 14:51:17 +03:30
patternihaandClaude Opus 5.5 2db099b34b TUN inbound: Block DNS and IPv6 leaks outside the TUN on Windows; Add strictRoute
Windows sends name queries to the DNS servers of all interfaces, and a
resolver on the local network (e.g. 192.168.1.1 from DHCP) is reached
through its more specific LAN route instead of the TUN, so DNS leaks
past it. IPv6 bypasses a TUN that cannot carry it.

With autoSystemRoutingTable set, the Windows TUN now adds Windows
Filtering Platform filters, all in one transaction and in a dynamic
session, so that they are removed when Xray exits, even if it crashes:
- DNS (port 53) only goes through the TUN, in both directions: its local
  address, and the interface it leaves or arrives by, must be the TUN's.
- IPv6 is blocked in both directions when the TUN has no IPv6 address or
  no IPv6 route, except loopback, neighbor and multicast listener
  discovery, and DHCPv6.
- Xray's own traffic is exempt: its connections out with a hard permit,
  which Windows Firewall rules do not override (like sing-box's
  strict_route), connections to its inbounds with an ordinary one.
If the filters cannot be added, the TUN does not start on Windows 10 and
later (only a warning on 7/8). The new `strictRoute` option (true by
default) turns them off.

Also on Windows:
- A warning for `dns` servers outside gateway and autoSystemRoutingTable,
  as queries to them cannot go through the TUN and are blocked.
- While DNS is restricted and autoOutboundsInterface is in use, Xray
  resolves the names it would ask Windows for itself (Go's resolver on
  its own sockets). Those lookups and the `localhost` DNS server skip the
  TUN's DNS servers, unless another interface uses them too, instead of
  looping back into the TUN.
- The DNS cache is flushed when the TUN starts and stops, and DNS
  registration is turned off on the TUN (through netsh before Windows 10
  1809).
- Close no longer panics when registering the route or interface change
  callbacks failed.
The README's Windows section describes all of it.

Tested on Windows 11, elevated, amd64 and 386: the filters, DNS arriving
through a real Wintun adapter and blocked outside it, the IPv6 block,
Windows Firewall rules, and a real Xray run. Windows 7/8 and Windows 10
before 1809 are untested.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-28 08:15:40 +03:30
30 changed files with 2014 additions and 1442 deletions
+3
View File
@@ -97,6 +97,9 @@ func New() *Client {
r := &net.Resolver{
PreferGo: true,
Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
if internet.IsSkippedDNSServer(address) {
return nil, errors.New("skipped DNS server ", address)
}
return d.DialContext(ctx, network, address)
},
}
+23
View File
@@ -0,0 +1,23 @@
package localdns
import (
"context"
"net/netip"
"testing"
"github.com/xtls/xray-core/transport/internet"
)
func TestSkippedDNSServers(t *testing.T) {
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("203.0.113.53")})
t.Cleanup(func() { internet.SkipDNSServers(nil) })
c := New()
if _, err := c.r.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil {
t.Error("a skipped DNS server was dialed")
}
conn, err := c.r.Dial(context.Background(), "udp", "127.0.0.1:53")
if err != nil {
t.Fatal(err)
}
conn.Close()
}
+32 -2
View File
@@ -5,8 +5,12 @@ import (
"fmt"
"math/big"
"net"
"runtime"
"slices"
"strconv"
"strings"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/proxy/tun"
"google.golang.org/protobuf/proto"
)
@@ -20,7 +24,8 @@ type TunConfig struct {
UserLevel uint32 `json:"userLevel"`
AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"`
AutoOutboundsInterface *string `json:"autoOutboundsInterface"`
AutoSystemDNS bool `json:"autoSystemDNS"`
AutoSystemDnsToGateway bool `json:"autoSystemDnsToGateway"`
AutoSystemWfpBlockLeak []string `json:"autoSystemWfpBlockLeak"`
}
func (v *TunConfig) Build() (proto.Message, error) {
@@ -32,7 +37,32 @@ func (v *TunConfig) Build() (proto.Message, error) {
DNS: v.DNS,
UserLevel: v.UserLevel,
AutoSystemRoutingTable: v.AutoSystemRoutingTable,
AutoSystemDns: v.AutoSystemDNS,
AutoSystemDnsToGateway: v.AutoSystemDnsToGateway,
}
for _, leak := range v.AutoSystemWfpBlockLeak {
switch leak := strings.ToLower(leak); leak {
case "dns", "misconfigtun":
config.AutoSystemWfpBlockLeak = append(config.AutoSystemWfpBlockLeak, leak)
default:
return nil, errors.New("unknown autoSystemWfpBlockLeak value: ", leak)
}
}
// Each option needs other settings on the system it takes effect on: the
// filters go along with the routes of autoSystemRoutingTable, "dns" lets
// DNS through the TUN only, and autoSystemDnsToGateway points the system
// DNS at the gateway.
switch runtime.GOOS {
case "windows":
if len(config.AutoSystemWfpBlockLeak) > 0 && len(v.AutoSystemRoutingTable) == 0 {
return nil, errors.New("autoSystemWfpBlockLeak needs autoSystemRoutingTable to be set")
}
if slices.Contains(config.AutoSystemWfpBlockLeak, "dns") && len(v.DNS) == 0 {
return nil, errors.New(`autoSystemWfpBlockLeak "dns" needs dns to be set`)
}
case "linux":
if v.AutoSystemDnsToGateway && len(v.Gateway) == 0 {
return nil, errors.New("autoSystemDnsToGateway needs gateway to be set")
}
}
if v.AutoOutboundsInterface != nil {
config.AutoOutboundsInterface = *v.AutoOutboundsInterface
+71
View File
@@ -0,0 +1,71 @@
package conf_test
import (
"encoding/json"
"runtime"
"testing"
. "github.com/xtls/xray-core/infra/conf"
"github.com/xtls/xray-core/proxy/tun"
)
func TestTunConfigAutoSystem(t *testing.T) {
creator := func() Buildable {
return new(TunConfig)
}
runMultiTestCase(t, []TestCase{
{
Input: `{"name": "xray0"}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500},
},
{
Input: `{"name": "xray0", "gateway": ["10.0.0.1/24"], "autoSystemDnsToGateway": true}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, Gateway: []string{"10.0.0.1/24"}, AutoSystemDnsToGateway: true},
},
{
Input: `{"name": "xray0", "dns": ["1.1.1.1"], "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["dns", "misconfigtun"]}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, DNS: []string{"1.1.1.1"}, AutoSystemRoutingTable: []string{"0.0.0.0/0"}, AutoOutboundsInterface: "auto", AutoSystemWfpBlockLeak: []string{"dns", "misconfigtun"}},
},
{
Input: `{"name": "xray0", "dns": ["1.1.1.1"], "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["DNS"]}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, DNS: []string{"1.1.1.1"}, AutoSystemRoutingTable: []string{"0.0.0.0/0"}, AutoOutboundsInterface: "auto", AutoSystemWfpBlockLeak: []string{"dns"}},
},
})
}
// TestTunConfigAutoSystemNeeds checks that an option is rejected without the
// setting it needs, only on the system it takes effect on.
func TestTunConfigAutoSystemNeeds(t *testing.T) {
for _, c := range []struct {
input string
goos string // where it is rejected
}{
{`{"name": "xray0", "autoSystemWfpBlockLeak": ["misconfigtun"]}`, "windows"},
{`{"name": "xray0", "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["misconfigtun"]}`, ""},
{`{"name": "xray0", "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["dns"]}`, "windows"},
{`{"name": "xray0", "autoSystemDnsToGateway": true}`, "linux"},
} {
config := new(TunConfig)
if err := json.Unmarshal([]byte(c.input), config); err != nil {
t.Fatal(err)
}
if _, err := config.Build(); (err != nil) != (runtime.GOOS == c.goos) {
t.Errorf("%s on %s: error = %v", c.input, runtime.GOOS, err)
}
}
}
func TestTunConfigAutoSystemWfpBlockLeakUnknown(t *testing.T) {
config := new(TunConfig)
if err := json.Unmarshal([]byte(`{"name": "xray0", "autoSystemWfpBlockLeak": ["dns", "ip"]}`), config); err != nil {
t.Fatal(err)
}
if _, err := config.Build(); err == nil {
t.Error("an unknown autoSystemWfpBlockLeak value was accepted")
}
}
-15
View File
@@ -277,7 +277,6 @@ func (w *VisionReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
w.ob.CanSpliceCopy = 1
}
}
SuppressOuterCloseNotify(w.conn)
readerConn, readCounter, _ := UnwrapRawConn(w.conn)
w.directReadCounter = readCounter
w.Reader = buf.NewReader(readerConn)
@@ -341,7 +340,6 @@ func (w *VisionWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
// w.ob.CanSpliceCopy = 1
// }
}
SuppressOuterCloseNotify(w.conn)
rawConn, _, writerCounter := UnwrapRawConn(w.conn)
w.Writer = buf.NewWriter(rawConn)
w.directWriteCounter = writerCounter
@@ -671,19 +669,6 @@ func XtlsFilterTls(buffer buf.MultiBuffer, trafficState *TrafficState, ctx conte
}
}
type CloseNotifySuppressor interface {
SuppressCloseNotify()
}
// Close our local TLS conn instance might send a incorrect close_notify alert
// if the XTLS direct copy mode is enabled and cause TLS BAD_RECORD_MAC on users' browser
// Close the underlying connection directly to avoid this issue.
func SuppressOuterCloseNotify(conn net.Conn) {
if suppressor, ok := stat.TryUnwrapStatsConn(conn).(CloseNotifySuppressor); ok {
suppressor.SuppressCloseNotify()
}
}
// UnwrapRawConn support unwrap encryption, stats, mask wrappers, tls, utls, reality, proxyproto, uds-wrapper conn and get raw tcp/uds conn from it
func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) {
var readCounter, writerCounter stats.Counter
+118 -38
View File
@@ -2,6 +2,7 @@ package shadowsocks_2022
import (
"context"
"io"
"time"
"github.com/xtls/xray-core/common"
@@ -12,6 +13,9 @@ import (
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/features/routing"
@@ -97,29 +101,35 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
return errors.New("unable to set read deadline").Base(err)
}
// 1. Single read call for Salt + Fixed-length header chunk per SIP022 §3.1.4
headerLen := i.method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize
headerBuf := make([]byte, headerLen)
n, err := conn.Read(headerBuf)
if err != nil || n < headerLen {
ResetTCPConn(conn)
return errors.New("failed to read complete handshake header")
}
var salt [32]byte
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
saltSlice := salt[:i.method.KeySaltLength]
fixedChunk := headerBuf[i.method.KeySaltLength:]
reader, reqHeader, err := InitServerStream(conn, i.method, i.psk, saltSlice, salt, fixedChunk, i.saltFilter)
if err != nil {
ResetTCPConn(conn)
if _, err := io.ReadFull(conn, saltSlice); err != nil {
return err
}
if !i.saltFilter.Check(salt) {
return ErrSaltNotUnique
}
sessionKey := DeriveSessionSubKey(i.psk, saltSlice, i.method.KeySaltLength)
aead, err := i.method.NewAEAD(sessionKey)
if err != nil {
return err
}
reader := NewStreamReader(conn, aead)
reqHeader, err := ReadClientRequestHeader(conn, reader)
if err != nil {
return err
}
conn.SetReadDeadline(time.Time{})
dest := reqHeader.Destination
writer := NewServerStreamWriter(conn, i.method, i.psk, saltSlice)
writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil)
if err != nil {
return err
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(),
@@ -136,17 +146,42 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
}
if len(reqHeader.EarlyData) > 0 {
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
earlyBuf := buf.New()
earlyBuf.Write(reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
return err
}
}
return TransportTCP(ctx, i.policyManager.ForLevel(uint32(i.user.Level)), reader, writer, link)
sessionPolicy = i.policyManager.ForLevel(uint32(i.user.Level))
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
}
func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn)
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
@@ -156,30 +191,75 @@ func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatch
for _, b := range mb {
decoded, err := i.udpCodec.DecodePacket(b.Bytes())
b.Release()
if err != nil || decoded.HeaderType != HeaderTypeClient {
continue
}
sessionItem := i.udpCodec.GetSession(decoded.SessionID)
if sessionItem.User == nil {
sessionItem.Lock()
if sessionItem.User == nil {
sessionItem.User = i.user
}
sessionItem.Unlock()
}
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(dest net.Destination, payload []byte) ([]byte, error) {
return i.udpCodec.EncodeServerPacket(decoded.SessionID, dest, payload)
})
if err != nil {
b.Release()
continue
}
entry, ok := udpConns.Load(decoded.SessionID)
if !ok {
sessCtx, cancel := context.WithCancel(ctx)
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: decoded.Destination,
Status: log.AccessAccepted,
Email: i.user.Email,
})
link, err := dispatcher.Dispatch(sessCtx, decoded.Destination)
if err != nil {
cancel()
b.Release()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(uint32(i.user.Level))
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(decoded.SessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(decoded.SessionID, newEntry)
if loaded {
// Another goroutine/packet beat us to storing, terminate our redundant link
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(sessID uint64, dest net.Destination, cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
encPacket, err := i.udpCodec.EncodeServerPacket(sessID, dest, rb.Bytes())
rb.Release()
if err != nil {
continue
}
_, _ = conn.Write(encPacket)
}
}
}(decoded.SessionID, decoded.Destination, entry)
}
}
entry.timer.Update()
payloadBuf := buf.New()
payloadBuf.Write(decoded.Payload)
payloadBuf.UDP = &decoded.Destination
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
b.Release()
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
}
}
}
+201 -48
View File
@@ -4,6 +4,7 @@ import (
"context"
"crypto/cipher"
"encoding/binary"
"io"
"strconv"
"strings"
"sync"
@@ -18,6 +19,8 @@ import (
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/common/uuid"
"github.com/xtls/xray-core/core"
@@ -204,46 +207,64 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
return errors.New("unable to set read deadline").Base(err)
}
// 1. Single read call for Salt + EIH + Fixed-length header chunk per SIP022 §3.1.4
headerLen := i.method.KeySaltLength + AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize
headerBuf := make([]byte, headerLen)
n, err := conn.Read(headerBuf)
if err != nil || n < headerLen {
ResetTCPConn(conn)
return errors.New("failed to read complete handshake header")
}
// 1. Read Request Salt (16 or 32 bytes)
var salt [32]byte
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
saltSlice := salt[:i.method.KeySaltLength]
eih := headerBuf[i.method.KeySaltLength : i.method.KeySaltLength+AESBlockSize]
fixedChunk := headerBuf[i.method.KeySaltLength+AESBlockSize:]
decryptedHash, err := DecryptEIH(i.method, i.masterPSK, saltSlice, eih)
if err != nil {
ResetTCPConn(conn)
if _, err := io.ReadFull(conn, saltSlice); err != nil {
return err
}
if !i.saltFilter.Check(salt) {
return ErrSaltNotUnique
}
// 2. Read Extended Identity Header (16 bytes)
var eih [AESBlockSize]byte
if _, err := io.ReadFull(conn, eih[:]); err != nil {
return err
}
// Decrypt EIH with IdentitySubKey derived from masterPSK and salt
identitySubkey := DeriveIdentitySubKey(i.masterPSK, saltSlice, i.method.KeySaltLength)
block, err := i.method.NewBlock(identitySubkey)
if err != nil {
return err
}
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih[:])
// Lookup user
user, ok := i.usersByHash.Load(decryptedHash)
if !ok {
ResetTCPConn(conn)
if !ok || user == nil {
return ErrInvalidRequest
}
userPSK := user.Account.(*MemoryAccount).Key
reader, reqHeader, err := InitServerStream(conn, i.method, userPSK, saltSlice, salt, fixedChunk, i.saltFilter)
// 3. Derive Session Subkey using matched user's PSK
sessionKey := DeriveSessionSubKey(userPSK, saltSlice, i.method.KeySaltLength)
aead, err := i.method.NewAEAD(sessionKey)
if err != nil {
ResetTCPConn(conn)
return err
}
reader := NewStreamReader(conn, aead)
// 4 & 5. Read Client Request Header
reqHeader, err := ReadClientRequestHeader(conn, reader)
if err != nil {
return err
}
conn.SetReadDeadline(time.Time{})
dest := reqHeader.Destination
writer := NewServerStreamWriter(conn, i.method, userPSK, saltSlice)
// 6. Send Server Response Handshake
writer, err := WriteTCPResponse(conn, i.method, userPSK, saltSlice, nil)
if err != nil {
return err
}
// Dispatch Connection to Xray routing with matched User
// 7. Dispatch Connection to Xray routing with matched User
inbound := session.InboundFromContext(ctx)
inbound.User = user
@@ -262,17 +283,42 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
}
if len(reqHeader.EarlyData) > 0 {
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
earlyBuf := buf.New()
earlyBuf.Write(reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
return err
}
}
return TransportTCP(ctx, i.policyManager.ForLevel(user.Level), reader, writer, link)
sessionPolicy = i.policyManager.ForLevel(user.Level)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
}
func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn)
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
@@ -296,61 +342,168 @@ func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection,
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
// Replay protection & session lookup
sessionItem := i.udpSessions.GetOrCreate(sessionID)
if !sessionItem.CheckPacketID(packetID) {
sessionItem.Lock()
if !sessionItem.Window.Check(packetID) {
sessionItem.Unlock()
b.Release()
continue
}
var userPSK []byte
var currentUser *protocol.MemoryUser
sessionItem.Lock()
currentUser = sessionItem.User
userPSK = sessionItem.UserPSK
sessionItem.Unlock()
if currentUser == nil {
if sessionItem.User != nil {
currentUser = sessionItem.User
userPSK = sessionItem.UserPSK
sessionItem.Unlock()
} else {
sessionItem.Unlock()
// Decrypt EIH
decryptedHash := DecryptUDPEIH(i.udpMasterCipher, rawHeader[:], packetBytes[16:32])
identitySubkey := DeriveIdentitySubKey(i.masterPSK, rawHeader[:8], i.method.KeySaltLength)
idBlock, err := i.method.NewBlock(identitySubkey)
if err != nil {
b.Release()
continue
}
var decryptedHash [16]byte
idBlock.Decrypt(decryptedHash[:], packetBytes[16:32])
user, ok := i.usersByHash.Load(decryptedHash)
if !ok {
if !ok || user == nil {
b.Release()
continue
}
currentUser = user
userPSK = user.Account.(*MemoryAccount).Key
sessionItem.Lock()
sessionItem.User = user
sessionItem.UserPSK = userPSK
sessionItem.Unlock()
}
decoded, err := sessionItem.DecryptAESPayload(i.method, userPSK, sessionID, packetID, rawHeader[:], packetBytes[32:])
// Decrypt Body (with AEAD caching per session)
bodyAead := sessionItem.GetRemoteCipher()
if bodyAead == nil {
bodyKey := DeriveSessionSubKey(userPSK, rawHeader[:8], i.method.KeySaltLength)
var err error
bodyAead, err = i.method.NewAEAD(bodyKey)
if err != nil {
b.Release()
continue
}
sessionItem.SetRemoteCipher(bodyAead)
}
bodyNonce := rawHeader[4:16]
bodyCipher := packetBytes[32:]
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
b.Release()
if err != nil {
if err != nil || len(bodyPlain) < 1+8+2 {
continue
}
sessionItem.Lock()
if sessionItem.User == nil {
sessionItem.User = currentUser
sessionItem.UserPSK = userPSK
}
sessionItem.Window.Add(packetID)
sessionItem.Unlock()
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(replyDest net.Destination, payload []byte) ([]byte, error) {
return i.encodeServerUDPPacket(sessionID, userPSK, replyDest, payload)
})
if bodyPlain[0] != HeaderTypeClient {
continue
}
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
diff := time.Now().Unix() - int64(epoch)
if diff < -30 || diff > 30 {
continue
}
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[9:11]))
offset := 11 + paddingLen
if len(bodyPlain) < offset {
continue
}
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
if err != nil {
continue
}
payload := bodyPlain[offset+addrLen:]
entry, ok := udpConns.Load(sessionID)
if !ok {
sessCtx, cancel := context.WithCancel(ctx)
inbound := session.InboundFromContext(sessCtx)
inbound.User = currentUser
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: currentUser.Email,
})
link, err := dispatcher.Dispatch(sessCtx, dest)
if err != nil {
cancel()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(currentUser.Level)
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(sessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
if loaded {
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(sessID uint64, uPSK []byte, d net.Destination, cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
encPacket, err := i.encodeServerUDPPacket(sessID, uPSK, d, rb.Bytes())
rb.Release()
if err != nil {
continue
}
_, _ = conn.Write(encPacket)
}
}
}(sessionID, userPSK, dest, entry)
}
}
entry.timer.Update()
pBuf := buf.New()
pBuf.Write(decoded.Payload)
pBuf.UDP = &decoded.Destination
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
pBuf.Write(payload)
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
}
}
}
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
return i.udpSessions.EncodeServerPacket(i.method, userPSK, clientSessionID, dest, payload)
sessionItem := i.udpSessions.GetOrCreate(clientSessionID)
if err := sessionItem.EnsureServerState(i.method, i.udpMasterCipher, nil, userPSK); err != nil {
return nil, err
}
return sessionItem.EncodeServerPacket(i.method, clientSessionID, dest, payload)
}
+124 -59
View File
@@ -4,6 +4,7 @@ import (
"context"
"crypto/cipher"
"encoding/binary"
"io"
"strconv"
"time"
@@ -14,6 +15,9 @@ import (
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/common/uuid"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
@@ -31,17 +35,18 @@ type relayDest struct {
destination net.Destination
email string
level uint32
key []byte
blockCipher cipher.Block
}
type RelayInbound struct {
networks []net.Network
method *CipherMethod
relayPSK []byte
relayBlock cipher.Block
destinations map[[AESBlockSize]byte]*relayDest
udpSessions *UDPSessionManager
policyManager policy.Manager
networks []net.Network
method *CipherMethod
relayPSK []byte
relayBlock cipher.Block
destinations map[[AESBlockSize]byte]*relayDest
rawDestinations []*RelayDestination
policyManager policy.Manager
}
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
@@ -73,13 +78,13 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
v := core.MustFromContext(ctx)
i := &RelayInbound{
networks: networks,
method: method,
relayPSK: relayPSK,
relayBlock: relayBlock,
destinations: make(map[[AESBlockSize]byte]*relayDest),
udpSessions: NewUDPSessionManager(500 * time.Second),
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
networks: networks,
method: method,
relayPSK: relayPSK,
relayBlock: relayBlock,
destinations: make(map[[AESBlockSize]byte]*relayDest),
rawDestinations: config.Destinations,
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
}
for idx, d := range config.Destinations {
@@ -103,6 +108,7 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)),
email: d.Email,
level: uint32(d.Level),
key: destKey,
blockCipher: destBlock,
}
}
@@ -133,36 +139,28 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
return errors.New("unable to set read deadline").Base(err)
}
// Read initial handshake in a single read call per SIP022 §3.1.3 & §3.1.4
// Read Salt + Outer EIH
needed := i.method.KeySaltLength + AESBlockSize
requestHeader := buf.New()
n, err := requestHeader.ReadFrom(conn)
if err != nil {
requestHeader.Release()
ResetTCPConn(conn)
var headerBuf [48]byte
headerSlice := headerBuf[:needed]
if _, err := io.ReadFull(conn, headerSlice); err != nil {
return err
}
if int(n) < needed {
requestHeader.Release()
ResetTCPConn(conn)
return ErrInvalidRequest
}
headerSlice := requestHeader.Bytes()
salt := headerSlice[:i.method.KeySaltLength]
eih := headerSlice[i.method.KeySaltLength:needed]
eih := headerSlice[i.method.KeySaltLength:]
decryptedHash, err := DecryptEIH(i.method, i.relayPSK, salt, eih)
identitySubkey := DeriveIdentitySubKey(i.relayPSK, salt, i.method.KeySaltLength)
block, err := i.method.NewBlock(identitySubkey)
if err != nil {
requestHeader.Release()
ResetTCPConn(conn)
return err
}
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih)
targetDest, ok := i.destinations[decryptedHash]
if !ok {
requestHeader.Release()
ResetTCPConn(conn)
return ErrInvalidRequest
}
conn.SetReadDeadline(time.Time{})
@@ -184,26 +182,45 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
link, err := dispatcher.Dispatch(ctx, targetDest.destination)
if err != nil {
requestHeader.Release()
return err
}
// Unwrap outer EIH: send client salt and remaining handshake bytes to next hop
// in a single write call, satisfying downstream server's single-read handshake expectation (SIP022 §3.1.3).
var saltCopy [32]byte
copy(saltCopy[:i.method.KeySaltLength], salt)
copy(requestHeader.Bytes()[AESBlockSize:AESBlockSize+i.method.KeySaltLength], saltCopy[:i.method.KeySaltLength])
requestHeader.Advance(AESBlockSize)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{requestHeader}); err != nil {
// Unwrap outer EIH: send client salt to next hop, stripping this hop's EIH
saltBuf := buf.New()
saltBuf.Write(salt)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{saltBuf}); err != nil {
return err
}
return TransportTCP(ctx, i.policyManager.ForLevel(targetDest.level), buf.NewReader(conn), buf.NewWriter(conn), link)
sessionPolicy = i.policyManager.ForLevel(targetDest.level)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(buf.NewReader(conn), link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(link.Reader, buf.NewWriter(conn), buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
}
func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn)
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
@@ -221,7 +238,11 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
var packetHeader [AESBlockSize]byte
i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize])
eiHeader := DecryptUDPEIH(i.relayBlock, packetHeader[:], data[AESBlockSize:2*AESBlockSize])
var eiHeader [AESBlockSize]byte
i.relayBlock.Decrypt(eiHeader[:], data[AESBlockSize:2*AESBlockSize])
for idx := 0; idx < AESBlockSize; idx++ {
eiHeader[idx] ^= packetHeader[idx]
}
targetDest, ok := i.destinations[eiHeader]
if !ok {
@@ -242,24 +263,68 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
dest := targetDest.destination
dest.Network = net.Network_UDP
sessionItem := i.udpSessions.GetOrCreate(sessionID)
if sessionItem.User == nil {
sessionItem.Lock()
if sessionItem.User == nil {
sessionItem.User = &protocol.MemoryUser{
Email: targetDest.email,
Level: targetDest.level,
}
entry, ok := udpConns.Load(sessionID)
if !ok {
sessCtx, cancel := context.WithCancel(ctx)
inbound := session.InboundFromContext(sessCtx)
inbound.User = &protocol.MemoryUser{
Email: targetDest.email,
Level: targetDest.level,
}
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: targetDest.email,
})
link, err := dispatcher.Dispatch(sessCtx, dest)
if err != nil {
cancel()
b.Release()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(targetDest.level)
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(sessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
if loaded {
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
_, _ = conn.Write(rb.Bytes())
rb.Release()
}
}
}(entry)
}
sessionItem.Unlock()
}
link, err := sessionItem.EnsureLink(ctx, conn, dest, dispatcher, i.policyManager, nil)
if err != nil {
b.Release()
continue
}
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
entry.timer.Update()
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
}
}
}
-11
View File
@@ -61,14 +61,3 @@ func DeriveUserPSKHash(userPSK []byte) [AESBlockSize]byte {
copy(out[:], h[:AESBlockSize])
return out
}
func DecryptEIH(method *CipherMethod, key, salt, eih []byte) ([AESBlockSize]byte, error) {
identitySubkey := DeriveIdentitySubKey(key, salt, method.KeySaltLength)
block, err := method.NewBlock(identitySubkey)
if err != nil {
return [AESBlockSize]byte{}, err
}
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih)
return decryptedHash, nil
}
+13 -33
View File
@@ -4,6 +4,7 @@ import (
"context"
"crypto/rand"
"io"
"time"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
@@ -45,12 +46,8 @@ func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) {
return nil, errors.New("invalid key: ", config.Key).Base(err)
}
if method.IsChaCha && len(pskList) > 1 {
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
}
finalPSK := pskList[len(pskList)-1]
udpCodec, err := NewUDPPacketCodec(method, pskList)
udpCodec, err := NewUDPPacketCodec(method, finalPSK)
if err != nil {
return nil, errors.New("failed to create udp packet codec").Base(err)
}
@@ -129,30 +126,18 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
var initialPayload []byte
var firstBuf *buf.Buffer
var remainingMB buf.MultiBuffer
if timeoutReader, ok := link.Reader.(buf.TimeoutReader); ok {
if mb, err := timeoutReader.ReadMultiBufferTimeout(0); err == nil && !mb.IsEmpty() {
remainingMB, firstBuf = buf.SplitFirst(mb)
initialPayload = firstBuf.Bytes()
}
}
bodyWriter, err := WriteTCPRequest(conn, o.method, o.pskList, destination, clientSaltSlice, initialPayload)
if firstBuf != nil {
firstBuf.Release()
}
bufferedWriter := buf.NewBufferedWriter(buf.NewWriter(conn))
bodyWriter, err := WriteTCPRequest(bufferedWriter, o.method, o.pskList, destination, clientSaltSlice, nil)
if err != nil {
buf.ReleaseMulti(remainingMB)
return errors.New("failed to write request").Base(err)
}
if !remainingMB.IsEmpty() {
if err := bodyWriter.WriteMultiBuffer(remainingMB); err != nil {
return err
}
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
return errors.New("failed to write A request payload").Base(err)
}
if err := bufferedWriter.SetBuffered(false); err != nil {
return err
}
return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer))
@@ -178,18 +163,13 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
}
if network == net.Network_UDP {
session, err := o.udpCodec.NewClientSession()
if err != nil {
return errors.New("failed to create client udp session").Base(err)
}
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
writer := &UDPWriter{
Writer: conn,
Destination: destination,
Session: session,
Codec: o.udpCodec,
}
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
@@ -202,8 +182,8 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
reader := &UDPReader{
Reader: conn,
Session: session,
Reader: conn,
Codec: o.udpCodec,
}
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
+187 -441
View File
@@ -16,13 +16,14 @@ import (
)
type UDPCodec struct {
method *CipherMethod
pskList [][]byte
psk []byte
blockCipher cipher.Block
blockCiphers []cipher.Block
chachaCipher cipher.AEAD
sessions *UDPSessionManager
method *CipherMethod
psk []byte
blockCipher cipher.Block
chachaCipher cipher.AEAD
clientBodyCipher cipher.AEAD
clientSessionID uint64
nextPacketID atomic.Uint64
sessions *UDPSessionManager
}
type (
@@ -47,23 +48,22 @@ func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
return c, nil
}
func NewUDPPacketCodec(method *CipherMethod, pskList [][]byte) (*UDPCodec, error) {
if method.IsChaCha && len(pskList) > 1 {
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
}
finalPSK := pskList[len(pskList)-1]
c, err := newUDPCodec(method, finalPSK)
func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
c, err := newUDPCodec(method, psk)
if err != nil {
return nil, err
}
c.pskList = pskList
if len(pskList) > 1 {
c.blockCiphers = make([]cipher.Block, len(pskList))
for i, psk := range pskList {
c.blockCiphers[i], err = method.NewBlock(psk)
if err != nil {
return nil, err
}
var sessID [8]byte
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
return nil, err
}
c.clientSessionID = binary.BigEndian.Uint64(sessID[:])
if !method.IsChaCha {
clientBodyKey := DeriveSessionSubKey(psk, sessID[:], method.KeySaltLength)
c.clientBodyCipher, err = method.NewAEAD(clientBodyKey)
if err != nil {
return nil, err
}
}
return c, nil
@@ -78,37 +78,108 @@ func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Dur
return c, nil
}
func (c *UDPCodec) Sessions() *UDPSessionManager {
return c.sessions
}
func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
packetID := c.nextPacketID.Add(1)
sessID := c.clientSessionID
func (c *UDPCodec) GetSession(sessionID uint64) *ServerUDPSession {
if c.sessions == nil {
return nil
// Padding determination (e.g. DNS port 53 disguise)
var paddingLen int
if dest.Port == 53 && len(payload) < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength-len(payload)) + 1
}
return c.sessions.GetOrCreate(sessionID)
addrPortLen := AddrPortLength(dest)
if c.method.IsChaCha {
// ChaCha20 mode: 24-byte nonce + plaintext header (27B) + padding + dest + payload + AEAD tag (16B)
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
var nonce [PacketNonceSize]byte
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(nonce[:])
var hdr [16 + 1 + 8 + 2]byte
binary.BigEndian.PutUint64(hdr[0:8], sessID)
binary.BigEndian.PutUint64(hdr[8:16], packetID)
hdr[16] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
plainBytes := outBuf.Bytes()[PacketNonceSize:]
outBuf.Extend(int32(c.chachaCipher.Overhead()))
c.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
return outBuf, nil
}
// AES mode:
// 16B Encrypted Header + (11B header + padding + dest + payload + 16B AEAD tag)
totalLen := 16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
var rawHeader [16]byte
binary.BigEndian.PutUint64(rawHeader[:8], sessID)
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
var encryptedHeader [16]byte
c.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
outBuf.Write(encryptedHeader[:])
bodyAead := c.clientBodyCipher
var hdr [1 + 8 + 2]byte
hdr[0] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
plainBytes := outBuf.Bytes()[16:]
bodyNonce := rawHeader[4:16]
outBuf.Extend(int32(bodyAead.Overhead()))
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
return outBuf, nil
}
type DecodedUDPPacket struct {
SessionID uint64
PacketID uint64
HeaderType byte
Timestamp uint64
ClientSessionID uint64
Destination net.Destination
Payload []byte
SessionID uint64
PacketID uint64
HeaderType byte
Timestamp uint64
Destination net.Destination
Payload []byte
}
func DecryptUDPEIH(block cipher.Block, rawHeader, eih []byte) [AESBlockSize]byte {
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih)
for k := 0; k < AESBlockSize; k++ {
decryptedHash[k] ^= rawHeader[k]
}
return decryptedHash
}
func ParseAddressPort(data []byte) (net.Destination, int, error) {
func parseAddressPort(data []byte) (net.Destination, int, error) {
if len(data) < 1 {
return net.Destination{}, 0, ErrPacketTooShort
}
@@ -149,9 +220,6 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
}
headerType := bodyPlain[0]
if headerType != HeaderTypeClient && headerType != HeaderTypeServer {
return DecodedUDPPacket{}, ErrBadHeaderType
}
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
if diff > 30 {
@@ -159,13 +227,11 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
}
offset := 9
var clientSessionID uint64
if headerType == HeaderTypeServer {
if len(bodyPlain) < offset+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort
}
clientSessionID = binary.BigEndian.Uint64(bodyPlain[offset : offset+8])
offset += 8
offset += 8 // skip clientSessionID
}
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
@@ -176,20 +242,19 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
}
offset += paddingLen
dest, addrLen, err := ParseAddressPort(bodyPlain[offset:])
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
if err != nil {
return DecodedUDPPacket{}, err
}
payload := bodyPlain[offset+addrLen:]
return DecodedUDPPacket{
SessionID: sessionID,
PacketID: packetID,
HeaderType: headerType,
Timestamp: epoch,
ClientSessionID: clientSessionID,
Destination: dest,
Payload: payload,
SessionID: sessionID,
PacketID: packetID,
HeaderType: headerType,
Timestamp: epoch,
Destination: dest,
Payload: payload,
}, nil
}
@@ -204,7 +269,7 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
}
nonce := data[:PacketNonceSize]
ciphertext := data[PacketNonceSize:]
plain, err := c.chachaCipher.Open(nil, nonce, ciphertext, nil)
plain, err := c.chachaCipher.Open(ciphertext[:0], nonce, ciphertext, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
}
@@ -215,22 +280,17 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
sessionID := binary.BigEndian.Uint64(plain[:8])
packetID := binary.BigEndian.Uint64(plain[8:16])
sessionItem := c.sessions.GetOrCreate(sessionID)
if !sessionItem.CheckPacketID(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
if c.sessions != nil {
sessionItem := c.sessions.GetOrCreate(sessionID)
sessionItem.Lock()
if !sessionItem.Window.CheckAndAdd(packetID) {
sessionItem.Unlock()
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
sessionItem.Unlock()
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
if err != nil {
return DecodedUDPPacket{}, err
}
if decoded.HeaderType != HeaderTypeClient {
return DecodedUDPPacket{}, ErrBadHeaderType
}
sessionItem.AddPacketID(packetID)
return decoded, nil
return parsePlainUDPPacket(sessionID, packetID, plain[16:])
}
// AES mode
@@ -239,52 +299,54 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
sessionItem := c.sessions.GetOrCreate(sessionID)
if !sessionItem.CheckPacketID(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
var bodyAead cipher.AEAD
var sessionItem *ServerUDPSession
return sessionItem.DecryptAESPayload(c.method, c.psk, sessionID, packetID, rawHeader[:], data[16:])
}
if c.sessions != nil {
sessionItem = c.sessions.GetOrCreate(sessionID)
sessionItem.Lock()
if !sessionItem.Window.Check(packetID) {
sessionItem.Unlock()
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
sessionItem.Unlock()
func (s *ServerUDPSession) DecryptAESPayload(method *CipherMethod, psk []byte, sessionID, packetID uint64, rawHeader, bodyCipher []byte) (DecodedUDPPacket, error) {
bodyAead := s.clientBodyCipher
isNewCipher := false
if bodyAead == nil {
bodyKey := DeriveSessionSubKey(psk, rawHeader[:8], method.KeySaltLength)
bodyAead = sessionItem.GetRemoteCipher()
if bodyAead == nil {
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
var err error
bodyAead, err = c.method.NewAEAD(bodyKey)
if err != nil {
return DecodedUDPPacket{}, err
}
sessionItem.SetRemoteCipher(bodyAead)
}
} else {
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
var err error
bodyAead, err = method.NewAEAD(bodyKey)
bodyAead, err = c.method.NewAEAD(bodyKey)
if err != nil {
return DecodedUDPPacket{}, err
}
isNewCipher = true
}
bodyNonce := rawHeader[4:16]
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
bodyCipher := data[16:]
bodyPlain, err := bodyAead.Open(bodyCipher[:0], bodyNonce, bodyCipher, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
if err != nil {
return DecodedUDPPacket{}, err
if sessionItem != nil {
sessionItem.Lock()
sessionItem.Window.Add(packetID)
sessionItem.Unlock()
}
if decoded.HeaderType != HeaderTypeClient {
return DecodedUDPPacket{}, ErrBadHeaderType
}
s.AddPacketID(packetID)
if isNewCipher {
s.clientBodyCipher = bodyAead
}
return decoded, nil
return parsePlainUDPPacket(sessionID, packetID, bodyPlain)
}
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) error {
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock cipher.Block, chachaCipher cipher.AEAD, psk []byte) error {
s.Lock()
defer s.Unlock()
if s.ServerSessionID != 0 {
@@ -301,29 +363,23 @@ func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) e
}
}
if method.IsChaCha {
var err error
s.serverChaCha, err = method.NewUDPCipher(psk)
return err
}
var err error
s.serverHeaderBlock, err = method.NewBlock(psk)
if err != nil {
s.ServerSessionID = 0
return err
}
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
s.serverBodyCipher, err = method.NewAEAD(bodyKey)
if err != nil {
s.ServerSessionID = 0
return err
s.ServerChaCha = chachaCipher
} else {
s.ServerBlockCipher = headerBlock
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
bodyAead, err := method.NewAEAD(bodyKey)
if err != nil {
s.ServerSessionID = 0
return err
}
s.ServerCipher = bodyAead
}
return nil
}
func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
serverSessionID := s.ServerSessionID
serverPacketID := s.ServerPacketID.Add(1) - 1
serverPacketID := s.ServerPacketID.Add(1)
if method.IsChaCha {
var nonce [PacketNonceSize]byte
@@ -348,7 +404,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
}
plainBuf.Write(payload)
sealed := s.serverChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil)
sealed := s.ServerChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil)
res := make([]byte, PacketNonceSize+len(sealed))
copy(res[:PacketNonceSize], nonce[:])
copy(res[PacketNonceSize:], sealed)
@@ -361,7 +417,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID)
var encryptedHeader [16]byte
s.serverHeaderBlock.Encrypt(encryptedHeader[:], rawHeader[:])
s.ServerBlockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
bodyBuf := buf.New()
defer bodyBuf.Release()
@@ -379,7 +435,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
bodyBuf.Write(payload)
bodyNonce := rawHeader[4:16]
sealedBody := s.serverBodyCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
sealedBody := s.ServerCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
res := make([]byte, 16+len(sealedBody))
copy(res[:16], encryptedHeader[:])
@@ -388,327 +444,17 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
}
func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
return c.sessions.EncodeServerPacket(c.method, c.psk, clientSessionID, dest, payload)
}
type serverSessionState struct {
sessionID uint64
window *SlidingWindow
cipher cipher.AEAD
lastSeen atomic.Int64
}
func (st *serverSessionState) check(packetID uint64) bool {
if st.window == nil {
st.window = new(SlidingWindow)
}
return st.window.Check(packetID)
}
func (st *serverSessionState) add(packetID uint64) {
if st.window == nil {
st.window = new(SlidingWindow)
}
st.window.Add(packetID)
}
type ClientUDPSession struct {
codec *UDPCodec
clientSessionID uint64
nextPacketID atomic.Uint64
clientBodyCipher cipher.AEAD
current atomic.Pointer[serverSessionState]
old atomic.Pointer[serverSessionState]
}
func (c *UDPCodec) NewClientSession() (*ClientUDPSession, error) {
var sessID [8]byte
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
sessionItem := c.sessions.GetOrCreate(clientSessionID)
if err := sessionItem.EnsureServerState(c.method, c.blockCipher, c.chachaCipher, c.psk); err != nil {
return nil, err
}
clientSessionID := binary.BigEndian.Uint64(sessID[:])
var clientBodyCipher cipher.AEAD
var err error
if !c.method.IsChaCha {
finalPSK := c.psk
clientBodyKey := DeriveSessionSubKey(finalPSK, sessID[:], c.method.KeySaltLength)
clientBodyCipher, err = c.method.NewAEAD(clientBodyKey)
if err != nil {
return nil, err
}
}
return &ClientUDPSession{
codec: c,
clientSessionID: clientSessionID,
clientBodyCipher: clientBodyCipher,
}, nil
}
func (s *ClientUDPSession) getServerSession(sessionID uint64, now int64) (*serverSessionState, error) {
cur := s.current.Load()
if cur != nil && cur.sessionID == sessionID {
return cur, nil
}
old := s.old.Load()
if old != nil && old.sessionID == sessionID {
if now-old.lastSeen.Load() > 60 {
s.old.CompareAndSwap(old, nil)
return nil, errors.New("old server session expired")
}
return old, nil
}
// New server session:
// Spec §3.2.4: reject newer server sessions when the last packet received from the old session is less than 1 minute old.
if old != nil && now-old.lastSeen.Load() < 60 {
return nil, errors.New("newer server session rejected: old session is less than 1 minute old")
}
var bodyAead cipher.AEAD
if !s.codec.method.IsChaCha {
var sessBytes [8]byte
binary.BigEndian.PutUint64(sessBytes[:], sessionID)
bodyKey := DeriveSessionSubKey(s.codec.psk, sessBytes[:], s.codec.method.KeySaltLength)
var err error
bodyAead, err = s.codec.method.NewAEAD(bodyKey)
if err != nil {
return nil, err
}
}
newState := &serverSessionState{
sessionID: sessionID,
cipher: bodyAead,
}
newState.lastSeen.Store(now)
if cur == nil {
s.current.CompareAndSwap(nil, newState)
return s.current.Load(), nil
}
s.old.Store(cur)
s.current.Store(newState)
return newState, nil
}
func (s *ClientUDPSession) ClientSessionID() uint64 {
return s.clientSessionID
}
func (s *ClientUDPSession) EncodePacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
packetID := s.nextPacketID.Add(1) - 1
sessID := s.clientSessionID
var paddingLen int
if dest.Port == 53 && len(payload) < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength) + 1
}
addrPortLen := AddrPortLength(dest)
if s.codec.method.IsChaCha {
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
var nonce [PacketNonceSize]byte
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(nonce[:])
var hdr [16 + 1 + 8 + 2]byte
binary.BigEndian.PutUint64(hdr[0:8], sessID)
binary.BigEndian.PutUint64(hdr[8:16], packetID)
hdr[16] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
plainBytes := outBuf.Bytes()[PacketNonceSize:]
outBuf.Extend(int32(s.codec.chachaCipher.Overhead()))
s.codec.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
return outBuf, nil
}
// AES mode
var sessBytes [8]byte
binary.BigEndian.PutUint64(sessBytes[:], sessID)
var rawHeader [16]byte
copy(rawHeader[:8], sessBytes[:])
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
eihCount := 0
if len(s.codec.pskList) > 1 {
eihCount = len(s.codec.pskList) - 1
}
totalLen := 16 + eihCount*16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
if len(s.codec.pskList) > 1 {
var encryptedHeader [16]byte
s.codec.blockCiphers[0].Encrypt(encryptedHeader[:], rawHeader[:])
outBuf.Write(encryptedHeader[:])
for i := 0; i < len(s.codec.pskList)-1; i++ {
nextPSK := s.codec.pskList[i+1]
pskHash := DeriveUserPSKHash(nextPSK)
var eihPlain [16]byte
for k := 0; k < 16; k++ {
eihPlain[k] = pskHash[k] ^ rawHeader[k]
}
var encryptedEIH [16]byte
s.codec.blockCiphers[i].Encrypt(encryptedEIH[:], eihPlain[:])
outBuf.Write(encryptedEIH[:])
}
} else {
var encryptedHeader [16]byte
s.codec.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
outBuf.Write(encryptedHeader[:])
}
bodyAead := s.clientBodyCipher
var hdr [1 + 8 + 2]byte
hdr[0] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
headerOffset := 16 + eihCount*16
plainBytes := outBuf.Bytes()[headerOffset:]
bodyNonce := rawHeader[4:16]
outBuf.Extend(int32(bodyAead.Overhead()))
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
return outBuf, nil
}
func (s *ClientUDPSession) DecodePacket(data []byte) (DecodedUDPPacket, error) {
if len(data) < PacketMinimalHeaderSize {
return DecodedUDPPacket{}, ErrPacketTooShort
}
if s.codec.method.IsChaCha {
if len(data) < PacketNonceSize+AEADTagSize {
return DecodedUDPPacket{}, ErrPacketTooShort
}
nonce := data[:PacketNonceSize]
ciphertext := data[PacketNonceSize:]
plain, err := s.codec.chachaCipher.Open(nil, nonce, ciphertext, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
}
if len(plain) < 16+1+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort
}
sessionID := binary.BigEndian.Uint64(plain[:8])
packetID := binary.BigEndian.Uint64(plain[8:16])
now := time.Now().Unix()
st, err := s.getServerSession(sessionID, now)
if err != nil {
return DecodedUDPPacket{}, err
}
if !st.check(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
if err != nil {
return DecodedUDPPacket{}, err
}
if decoded.HeaderType != HeaderTypeServer {
return DecodedUDPPacket{}, ErrBadHeaderType
}
if decoded.ClientSessionID != s.clientSessionID {
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
}
st.add(packetID)
st.lastSeen.Store(now)
return decoded, nil
}
// AES mode
var rawHeader [16]byte
s.codec.blockCipher.Decrypt(rawHeader[:], data[:16])
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
now := time.Now().Unix()
st, err := s.getServerSession(sessionID, now)
if err != nil {
return DecodedUDPPacket{}, err
}
if !st.check(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
bodyAead := st.cipher
bodyNonce := rawHeader[4:16]
bodyCipher := data[16:]
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
if err != nil {
return DecodedUDPPacket{}, err
}
if decoded.HeaderType != HeaderTypeServer {
return DecodedUDPPacket{}, ErrBadHeaderType
}
if decoded.ClientSessionID != s.clientSessionID {
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
}
st.add(packetID)
st.lastSeen.Store(now)
return decoded, nil
return sessionItem.EncodeServerPacket(c.method, clientSessionID, dest, payload)
}
type UDPWriter struct {
Writer io.Writer
Destination net.Destination
Session *ClientUDPSession
Codec *UDPPacketCodec
}
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
@@ -722,7 +468,7 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
if b.UDP != nil {
dest = *b.UDP
}
pktBuf, err := w.Session.EncodePacket(dest, b.Bytes())
pktBuf, err := w.Codec.EncodeClientPacket(dest, b.Bytes())
b.Release()
if err != nil {
buf.ReleaseMulti(mb)
@@ -739,8 +485,8 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
}
type UDPReader struct {
Reader io.Reader
Session *ClientUDPSession
Reader io.Reader
Codec *UDPPacketCodec
}
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
@@ -752,7 +498,7 @@ func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
return nil, err
}
decoded, err := r.Session.DecodePacket(buffer.Bytes())
decoded, err := r.Codec.DecodePacket(buffer.Bytes())
if err != nil {
buffer.Release()
continue
-105
View File
@@ -2,11 +2,9 @@ package shadowsocks_2022_test
import (
"context"
"crypto/rand"
"encoding/base64"
"encoding/binary"
"errors"
"io"
gonet "net"
"sync"
"sync/atomic"
@@ -271,106 +269,3 @@ func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error {
}
return nil
}
func TestRelayTCPHandshakeForwarding(t *testing.T) {
methods := []string{MethodAES128GCM, MethodAES256GCM}
for _, methodName := range methods {
t.Run(methodName, func(t *testing.T) {
method, err := GetCipherMethod(methodName)
common.Must(err)
relayKey := make([]byte, method.KeySaltLength)
destKey := make([]byte, method.KeySaltLength)
_, _ = io.ReadFull(rand.Reader, relayKey)
_, _ = io.ReadFull(rand.Reader, destKey)
targetPort := uint32(54321)
relayConfig := &RelayServerConfig{
Method: methodName,
Key: base64.StdEncoding.EncodeToString(relayKey),
Destinations: []*RelayDestination{
{
Key: base64.StdEncoding.EncodeToString(destKey),
Address: net.NewIPOrDomain(net.LocalHostIP),
Port: targetPort,
Email: "test@xray.com",
},
},
}
testCtx := newTestContext()
inbound, err := NewRelayServer(testCtx, relayConfig)
common.Must(err)
targetDest := net.TCPDestination(net.LocalHostIP, net.Port(targetPort))
downstreamR, downstreamW := gonet.Pipe()
defer downstreamR.Close()
defer downstreamW.Close()
disp := &dummyDispatcher{
onDispatch: func(ctx context.Context, dest net.Destination) (*transport.Link, error) {
inLink := &transport.Link{
Reader: buf.NewReader(downstreamR),
Writer: &customWriter{
write: func(mb buf.MultiBuffer) error {
defer buf.ReleaseMulti(mb)
for _, b := range mb {
if _, err := downstreamW.Write(b.Bytes()); err != nil {
return err
}
}
return nil
},
},
}
return inLink, nil
},
}
clientConn, relayConn := gonet.Pipe()
defer clientConn.Close()
defer relayConn.Close()
go func() {
_ = inbound.Process(testCtx, net.Network_TCP, &dummyStatConn{Conn: relayConn}, disp)
}()
clientSalt := make([]byte, method.KeySaltLength)
_, _ = io.ReadFull(rand.Reader, clientSalt)
pskList := [][]byte{relayKey, destKey}
go func() {
_, err := WriteTCPRequest(clientConn, method, pskList, targetDest, clientSalt, []byte("relay payload"))
if err != nil {
t.Errorf("WriteTCPRequest failed: %v", err)
}
}()
// Downstream server must be able to read Salt + Fixed chunk in a single Read call!
headerLen := method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize
headerBuf := make([]byte, headerLen)
n, err := downstreamR.Read(headerBuf)
if err != nil {
t.Fatalf("downstream failed to read handshake: %v", err)
}
if n < headerLen {
t.Fatalf("downstream expected single read >= %d bytes, got %d", headerLen, n)
}
// Verify downstream can decode the fixed chunk and subsequent payload
sessionKey := DeriveSessionSubKey(destKey, headerBuf[:method.KeySaltLength], method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
common.Must(err)
reader := NewStreamReader(downstreamR, aead)
reqHeader, err := ReadClientRequestHeaderWithFixed(reader, headerBuf[method.KeySaltLength:])
if err != nil {
t.Fatalf("downstream failed to parse client request header: %v", err)
}
if string(reqHeader.EarlyData) != "relay payload" {
t.Fatalf("payload mismatch: expected 'relay payload', got '%s'", string(reqHeader.EarlyData))
}
})
}
}
+16 -41
View File
@@ -6,11 +6,8 @@ import (
"sync/atomic"
"time"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/transport"
)
const (
@@ -77,42 +74,30 @@ func (f *SlidingWindow) CheckAndAdd(counter uint64) bool {
type ServerUDPSession struct {
sync.Mutex
SessionID uint64
Window *SlidingWindow
User *protocol.MemoryUser
UserPSK []byte
LastActive atomic.Int64 // Unix timestamp in seconds
clientBodyCipher cipher.AEAD
SessionID uint64
RemoteCipher atomic.Pointer[cipher.AEAD]
Window SlidingWindow
User *protocol.MemoryUser
UserPSK []byte
LastActive atomic.Int64 // Unix timestamp in seconds
ServerSessionID uint64
ServerPacketID atomic.Uint64
serverBodyCipher cipher.AEAD
serverHeaderBlock cipher.Block
serverChaCha cipher.AEAD
manager *UDPSessionManager
link atomic.Pointer[transport.Link]
timer *signal.ActivityTimer
currentConn atomic.Value // stores stat.Connection
ServerCipher cipher.AEAD
ServerBlockCipher cipher.Block
ServerChaCha cipher.AEAD
}
func (s *ServerUDPSession) CheckPacketID(packetID uint64) bool {
s.Lock()
defer s.Unlock()
if s.Window == nil {
s.Window = new(SlidingWindow)
func (s *ServerUDPSession) GetRemoteCipher() cipher.AEAD {
ptr := s.RemoteCipher.Load()
if ptr == nil {
return nil
}
return s.Window.Check(packetID)
return *ptr
}
func (s *ServerUDPSession) AddPacketID(packetID uint64) {
s.Lock()
defer s.Unlock()
if s.Window == nil {
s.Window = new(SlidingWindow)
}
s.Window.Add(packetID)
func (s *ServerUDPSession) SetRemoteCipher(c cipher.AEAD) {
s.RemoteCipher.Store(&c)
}
type UDPSessionManager struct {
@@ -137,7 +122,6 @@ func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession {
s := &ServerUDPSession{
SessionID: sessionID,
manager: m,
}
s.LastActive.Store(now)
@@ -164,7 +148,6 @@ func (m *UDPSessionManager) cleanup(now int64) {
m.sessions.Range(func(k uint64, v *ServerUDPSession) bool {
if now-v.LastActive.Load() > timeoutSec {
m.sessions.Delete(k)
v.Close()
}
return true
})
@@ -173,11 +156,3 @@ func (m *UDPSessionManager) cleanup(now int64) {
func (m *UDPSessionManager) Delete(sessionID uint64) {
m.sessions.Delete(sessionID)
}
func (m *UDPSessionManager) EncodeServerPacket(method *CipherMethod, psk []byte, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
sessionItem := m.GetOrCreate(clientSessionID)
if err := sessionItem.EnsureServerState(method, psk); err != nil {
return nil, err
}
return sessionItem.EncodeServerPacket(method, clientSessionID, dest, payload)
}
+6 -149
View File
@@ -2,161 +2,18 @@ package shadowsocks_2022
import (
"context"
"sync"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/log"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/features/routing"
"github.com/xtls/xray-core/proxy"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet/stat"
)
func (s *ServerUDPSession) UpdateConn(conn stat.Connection) {
if s.currentConn.Load() == nil {
s.currentConn.Store(conn)
}
if s.timer != nil {
s.timer.Update()
}
}
func (s *ServerUDPSession) WriteToClient(b []byte) error {
connVal := s.currentConn.Load()
if connVal == nil {
return errors.New("client connection closed")
}
conn, ok := connVal.(stat.Connection)
if !ok || conn == nil {
return errors.New("client connection closed")
}
_, err := conn.Write(b)
return err
}
func (s *ServerUDPSession) Close() {
if s.timer != nil {
s.timer.SetTimeout(0)
}
if link := s.link.Load(); link != nil {
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
}
}
func (s *ServerUDPSession) EnsureLink(
ctx context.Context,
conn stat.Connection,
dest net.Destination,
dispatcher routing.Dispatcher,
policyManager policy.Manager,
responseEncoder func(dest net.Destination, payload []byte) ([]byte, error),
) (*transport.Link, error) {
s.UpdateConn(conn)
if link := s.link.Load(); link != nil {
return link, nil
}
s.Lock()
defer s.Unlock()
if link := s.link.Load(); link != nil {
return link, nil
}
sessCtx, cancel := context.WithCancel(ctx)
inbound := session.InboundFromContext(sessCtx)
if inbound != nil && s.User != nil {
inbound.User = s.User
}
var email string
var level uint32
if s.User != nil {
email = s.User.Email
level = s.User.Level
}
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: email,
})
link, err := dispatcher.Dispatch(sessCtx, dest)
if err != nil {
cancel()
return nil, err
}
s.link.Store(link)
sessionPolicy := policyManager.ForLevel(level)
s.timer = signal.CancelAfterInactivity(sessCtx, func() {
if s.manager != nil {
s.manager.Delete(s.SessionID)
}
s.Close()
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
go handleUDPResponse(s, link, dest, responseEncoder)
return link, nil
}
// ResetTCPConn sets SO_LINGER to 0 per SIP022 §3.1.4 to consistently send RST on close
// when handshake or header validation fails.
func ResetTCPConn(conn net.Conn) {
rawConn, _, _ := proxy.UnwrapRawConn(conn)
if tcpConn, ok := rawConn.(*net.TCPConn); ok {
_ = tcpConn.SetLinger(0)
}
}
func handleUDPResponse(s *ServerUDPSession, link *transport.Link, fallbackDest net.Destination, encode func(dest net.Destination, payload []byte) ([]byte, error)) {
defer func() {
if s.timer != nil {
s.timer.SetTimeout(0)
}
}()
for {
resMb, err := link.Reader.ReadMultiBuffer()
if err != nil {
return
}
if s.timer != nil {
s.timer.Update()
}
for i, rb := range resMb {
b := rb.Bytes()
if encode != nil {
replyDest := fallbackDest
if rb.UDP != nil {
replyDest = *rb.UDP
}
encPacket, err := encode(replyDest, b)
rb.Release()
if err != nil {
continue
}
if err := s.WriteToClient(encPacket); err != nil {
buf.ReleaseMulti(resMb[i+1:])
return
}
} else {
err := s.WriteToClient(b)
rb.Release()
if err != nil {
buf.ReleaseMulti(resMb[i+1:])
return
}
}
}
}
type udpConnEntry struct {
sync.Mutex
link *transport.Link
timer *signal.ActivityTimer
cancel context.CancelFunc
}
const (
+36 -172
View File
@@ -182,48 +182,57 @@ func TestTCPStream(t *testing.T) {
common.Must(err)
IncreaseNonce(reader.Nonce())
dest, addrLen, err := ParseAddressPort(plainVar)
vBuf := buf.New()
vBuf.Write(plainVar)
receivedDest, err = ReadAddressPort(vBuf)
common.Must(err)
receivedDest = net.TCPDestination(dest.Address, dest.Port)
plainVar = plainVar[addrLen:]
padLen := int(binary.BigEndian.Uint16(plainVar[:2]))
receivedPayload = plainVar[2+padLen:]
// Server sends response stream with receivedPayload as first payload
writer := NewServerStreamWriter(serverConn, method, rawKey, salt)
pBuf := buf.New()
pBuf.Write(receivedPayload)
_ = writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
// Skip padding
var padBytes [2]byte
_, _ = vBuf.Read(padBytes[:])
padLen := int(padBytes[0])<<8 | int(padBytes[1])
vBuf.Advance(int32(padLen))
// Read and echo additional stream data
receivedPayload = make([]byte, vBuf.Len())
copy(receivedPayload, vBuf.Bytes())
vBuf.Release()
// Server sends response handshake
serverSalt := make([]byte, method.KeySaltLength)
_, _ = rand.Read(serverSalt)
respKey := DeriveSessionSubKey(rawKey, serverSalt, method.KeySaltLength)
respAead, err := method.NewAEAD(respKey)
writer := NewStreamWriter(serverConn, respAead)
_, _ = serverConn.Write(serverSalt)
fixedResp := make([]byte, 1+8+method.KeySaltLength+2)
fixedResp[0] = HeaderTypeServer
binary.BigEndian.PutUint64(fixedResp[1:9], uint64(time.Now().Unix()))
copy(fixedResp[9:9+method.KeySaltLength], salt)
binary.BigEndian.PutUint16(fixedResp[9+method.KeySaltLength:11+method.KeySaltLength], 0)
fixedChunk := respAead.Seal(nil, writer.Nonce(), fixedResp, nil)
IncreaseNonce(writer.Nonce())
_, _ = serverConn.Write(fixedChunk)
// Echo stream data
mb, err := reader.ReadMultiBuffer()
common.Must(err)
_ = writer.WriteMultiBuffer(mb)
_ = writer.Close()
}()
// Client goroutine
go func() {
defer wg.Done()
clientSalt := make([]byte, method.KeySaltLength)
common.Must2(io.ReadFull(rand.Reader, clientSalt))
writer, err := WriteTCPRequest(clientConn, method, [][]byte{rawKey}, dest, clientSalt, testPayload)
clientSalt, writer, err := ClientHandshake(clientConn, method, [][]byte{rawKey}, dest, testPayload)
common.Must(err)
reader, err := ReadTCPResponse(clientConn, method, rawKey, clientSalt)
reader, _, err := ClientVerifyServerResponse(clientConn, method, rawKey, clientSalt)
common.Must(err)
// The first ReadMultiBuffer drains initialPayload from reader cache
mbInit, err := reader.ReadMultiBuffer()
common.Must(err)
if !bytes.Equal(mbInit[0].Bytes(), testPayload) {
t.Errorf("drained initial payload mismatch: got %s, want %s", mbInit[0].Bytes(), testPayload)
}
buf.ReleaseMulti(mbInit)
// Send additional stream data
streamData := []byte("stream chunk test")
_ = writer.WriteMultiBuffer(buf.MultiBuffer{buf.FromBytes(streamData)})
_ = writer.WriteChunk(streamData)
mb, err := reader.ReadMultiBuffer()
common.Must(err)
@@ -263,14 +272,12 @@ func TestUDPCodec(t *testing.T) {
psk := make([]byte, method.KeySaltLength)
_, _ = rand.Read(psk)
clientCodec, err := NewUDPPacketCodec(method, [][]byte{psk})
clientCodec, err := NewUDPPacketCodec(method, psk)
common.Must(err)
serverCodec, err := NewUDPServerCodec(method, psk, time.Minute)
common.Must(err)
session, err := clientCodec.NewClientSession()
common.Must(err)
pktBuf, err := session.EncodePacket(dest, payload)
pktBuf, err := clientCodec.EncodeClientPacket(dest, payload)
common.Must(err)
defer pktBuf.Release()
@@ -353,146 +360,3 @@ func TestMultiUserManager(t *testing.T) {
t.Fatal("user1 should have been removed")
}
}
func TestLargeStreamTransfer(t *testing.T) {
method, err := GetCipherMethod(MethodAES128GCM)
common.Must(err)
sessionKey := make([]byte, 16)
_, _ = rand.Read(sessionKey)
clientAead, err := method.NewAEAD(sessionKey)
common.Must(err)
serverAead, err := method.NewAEAD(sessionKey)
common.Must(err)
r, w := io.Pipe()
defer r.Close()
defer w.Close()
writer := NewStreamWriter(w, clientAead)
reader := NewStreamReader(r, serverAead)
const totalSize = 100 * 1024 // 100 KB
data := make([]byte, totalSize)
_, _ = rand.Read(data)
errCh := make(chan error, 1)
go func() {
// Write using Write (which splits by MaxPacketSize = 65535)
_, werr := writer.Write(data)
if werr != nil {
errCh <- werr
return
}
_ = w.Close()
errCh <- nil
}()
var received []byte
for {
mb, rerr := reader.ReadMultiBuffer()
if !mb.IsEmpty() {
for _, b := range mb {
received = append(received, b.Bytes()...)
}
buf.ReleaseMulti(mb)
}
if rerr != nil {
if rerr == io.EOF {
break
}
t.Fatalf("ReadMultiBuffer error: %v", rerr)
}
}
if werr := <-errCh; werr != nil {
t.Fatalf("writer error: %v", werr)
}
if len(received) != totalSize {
t.Fatalf("received size mismatch: got %d, want %d", len(received), totalSize)
}
if !bytes.Equal(received, data) {
t.Fatal("received data does not match sent data")
}
}
func TestClientUDPSessionMultiDestination(t *testing.T) {
for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} {
t.Run(methodName, func(t *testing.T) {
method, err := GetCipherMethod(methodName)
common.Must(err)
rawKey := make([]byte, method.KeySaltLength)
_, _ = rand.Read(rawKey)
clientCodec, err := NewUDPPacketCodec(method, [][]byte{rawKey})
common.Must(err)
serverCodec, err := NewUDPServerCodec(method, rawKey, time.Minute)
common.Must(err)
session, err := clientCodec.NewClientSession()
common.Must(err)
dest1 := net.UDPDestination(net.LocalHostIP, net.Port(53))
dest2 := net.UDPDestination(net.IPAddress([]byte{127, 0, 0, 2}), net.Port(53))
payload1 := []byte("query-google-dns")
payload2 := []byte("query-cloudflare-dns")
// Client sends to dest1 and dest2 using SAME session
pkt1, err := session.EncodePacket(dest1, payload1)
common.Must(err)
defer pkt1.Release()
pkt2, err := session.EncodePacket(dest2, payload2)
common.Must(err)
defer pkt2.Release()
// Server decodes both
dec1, err := serverCodec.DecodePacket(pkt1.Bytes())
common.Must(err)
dec2, err := serverCodec.DecodePacket(pkt2.Bytes())
common.Must(err)
if dec1.SessionID != session.ClientSessionID() || dec2.SessionID != session.ClientSessionID() {
t.Fatalf("both packets must share client session ID %d, got %d and %d", session.ClientSessionID(), dec1.SessionID, dec2.SessionID)
}
if dec1.Destination.String() != dest1.String() {
t.Fatalf("expected dest1 %s, got %s", dest1, dec1.Destination)
}
if dec2.Destination.String() != dest2.String() {
t.Fatalf("expected dest2 %s, got %s", dest2, dec2.Destination)
}
if !bytes.Equal(dec1.Payload, payload1) || !bytes.Equal(dec2.Payload, payload2) {
t.Fatal("payload mismatch")
}
// Server replies to dest1 and dest2
respPayload1 := []byte("reply-google-dns")
respPayload2 := []byte("reply-cloudflare-dns")
respPkt1, err := serverCodec.EncodeServerPacket(dec1.SessionID, dest1, respPayload1)
common.Must(err)
respPkt2, err := serverCodec.EncodeServerPacket(dec2.SessionID, dest2, respPayload2)
common.Must(err)
// Client decodes replies
clientDec1, err := session.DecodePacket(respPkt1)
common.Must(err)
if clientDec1.Destination.String() != dest1.String() {
t.Fatalf("expected client dec1 dest %s, got %s", dest1, clientDec1.Destination)
}
if !bytes.Equal(clientDec1.Payload, respPayload1) {
t.Fatal("reply payload 1 mismatch")
}
clientDec2, err := session.DecodePacket(respPkt2)
common.Must(err)
if clientDec2.Destination.String() != dest2.String() {
t.Fatalf("expected client dec2 dest %s, got %s", dest2, clientDec2.Destination)
}
if !bytes.Equal(clientDec2.Payload, respPayload2) {
t.Fatal("reply payload 2 mismatch")
}
})
}
}
+115 -235
View File
@@ -1,25 +1,18 @@
package shadowsocks_2022
import (
"context"
"crypto/cipher"
"crypto/rand"
"encoding/binary"
"io"
"math"
mrand "math/rand/v2"
"sync"
"time"
"github.com/xtls/xray-core/common/antireplay"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/transport"
)
var addrParser = protocol.NewAddressParser(
@@ -45,6 +38,15 @@ func WriteAddressPort(w io.Writer, dest net.Destination) error {
return addrParser.WriteAddressPort(w, dest.Address, dest.Port)
}
// ReadAddressPort reads a destination address and port in SOCKS5 format
func ReadAddressPort(r io.Reader) (net.Destination, error) {
addr, port, err := addrParser.ReadAddressPort(nil, r)
if err != nil {
return net.Destination{}, err
}
return net.TCPDestination(addr, port), nil
}
// AddrPortLength returns the serialized length of a destination in SOCKS5 format
func AddrPortLength(dest net.Destination) int {
switch dest.Address.Family() {
@@ -117,16 +119,8 @@ func (w *StreamWriter) Write(p []byte) (int, error) {
func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
defer buf.ReleaseMulti(mb)
for _, b := range mb {
p := b.Bytes()
for len(p) > 0 {
chunkSize := len(p)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
if err := w.WriteChunk(p[:chunkSize]); err != nil {
return err
}
p = p[chunkSize:]
if err := w.WriteChunk(b.Bytes()); err != nil {
return err
}
}
return nil
@@ -174,7 +168,7 @@ func (r *StreamReader) Read(p []byte) (int, error) {
IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 || payloadLen > MaxPacketSize {
if payloadLen == 0 {
return 0, ErrInvalidRequest
}
@@ -200,10 +194,11 @@ func (r *StreamReader) Read(p []byte) (int, error) {
func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
if r.cached > 0 {
mb := buf.MergeBytes(nil, r.buffer[r.offset:r.offset+r.cached])
b := buf.New()
b.Write(r.buffer[r.offset : r.offset+r.cached])
r.cached = 0
r.offset = 0
return mb, nil
return buf.MultiBuffer{b}, nil
}
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
@@ -217,7 +212,7 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 || payloadLen > MaxPacketSize {
if payloadLen == 0 {
return nil, ErrInvalidRequest
}
@@ -232,8 +227,9 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
}
IncreaseNonce(r.nonce[:])
mb := buf.MergeBytes(nil, decryptedPayload)
return mb, nil
b := buf.New()
b.Write(decryptedPayload)
return buf.MultiBuffer{b}, nil
}
type ClientRequestHeader struct {
@@ -241,8 +237,13 @@ type ClientRequestHeader struct {
EarlyData []byte
}
func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (*ClientRequestHeader, error) {
plainFixed, err := reader.cipher.Open(fixedChunk[:0], reader.Nonce(), fixedChunk, nil)
func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientRequestHeader, error) {
var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte
if _, err := io.ReadFull(conn, fixedBuf[:]); err != nil {
return nil, err
}
plainFixed, err := reader.cipher.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil)
if err != nil {
return nil, errors.New("failed to decrypt client request header").Base(err)
}
@@ -271,7 +272,7 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
} else {
varChunkCipher = make([]byte, needed)
}
if _, err := io.ReadFull(reader.reader, varChunkCipher); err != nil {
if _, err := io.ReadFull(conn, varChunkCipher); err != nil {
return nil, err
}
@@ -281,34 +282,31 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
}
IncreaseNonce(reader.Nonce())
dest, addrLen, err := ParseAddressPort(plainVar)
b := buf.New()
b.Write(plainVar)
defer b.Release()
dest, err := ReadAddressPort(b)
if err != nil {
return nil, err
}
dest.Network = net.Network_TCP
offset := addrLen
if len(plainVar) < offset+2 {
return nil, ErrPacketTooShort
var padLenBytes [2]byte
if _, err := b.Read(padLenBytes[:]); err != nil {
return nil, err
}
paddingLen := int(binary.BigEndian.Uint16(plainVar[offset : offset+2]))
offset += 2
if len(plainVar) < offset+paddingLen {
paddingLen := int(binary.BigEndian.Uint16(padLenBytes[:]))
if int(b.Len()) < paddingLen {
return nil, ErrNoPadding
}
offset += paddingLen
var earlyData []byte
var payloadLen int
if len(plainVar) > offset {
earlyData = plainVar[offset:]
payloadLen = len(earlyData)
if paddingLen > 0 {
b.Advance(int32(paddingLen))
}
// SIP022 §3.1.4: Servers MUST reject the request if the variable-length header chunk does not contain payload and the padding length is 0.
if paddingLen == 0 && payloadLen == 0 {
return nil, errors.New("request without payload and padding is not allowed")
var earlyData []byte
if b.Len() > 0 {
earlyData = make([]byte, b.Len())
copy(earlyData, b.Bytes())
}
return &ClientRequestHeader{
@@ -317,6 +315,34 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
}, nil
}
// ClientHandshake writes the full client request header to w
func ClientHandshake(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, payload []byte) ([]byte, *StreamWriter, error) {
salt := make([]byte, method.KeySaltLength)
if _, err := io.ReadFull(rand.Reader, salt); err != nil {
return nil, nil, err
}
writer, err := WriteTCPRequest(w, method, pskList, dest, salt, payload)
if err != nil {
return nil, nil, err
}
return salt, writer.(*StreamWriter), nil
}
// ClientVerifyServerResponse reads and verifies the server's handshake response
func ClientVerifyServerResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (*StreamReader, []byte, error) {
reader, err := ReadTCPResponse(r, method, psk, clientSalt)
if err != nil {
return nil, nil, err
}
sr := reader.(*StreamReader)
var initialPayload []byte
if sr.cached > 0 {
initialPayload = make([]byte, sr.cached)
copy(initialPayload, sr.buffer[sr.offset:sr.offset+sr.cached])
}
return sr, initialPayload, nil
}
// WriteTCPRequest writes the Shadowsocks 2022 request header into w and returns a body writer.
func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) {
finalPSK := pskList[len(pskList)-1]
@@ -328,16 +354,7 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
writer := NewStreamWriter(w, aead)
payloadLen := len(payload)
var paddingLen int
if payloadLen < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength) + 1
}
addrPortLen := AddrPortLength(dest)
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
totalHandshakeLen := int32(method.KeySaltLength + len(pskList)*AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize + varHeaderLen + AEADTagSize)
handshakeBuf := buf.NewWithSize(totalHandshakeLen)
handshakeBuf := buf.New()
defer handshakeBuf.Release()
handshakeBuf.Write(clientSalt)
@@ -355,6 +372,14 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
handshakeBuf.Write(encryptedEIH[:])
}
payloadLen := len(payload)
var paddingLen int
if payloadLen < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength-payloadLen) + 1
}
addrPortLen := AddrPortLength(dest)
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte
fixedHeaderPlaintext[0] = HeaderTypeClient
binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix()))
@@ -364,7 +389,7 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
IncreaseNonce(writer.nonce[:])
handshakeBuf.Write(fixedChunk)
varHeaderBuf := buf.NewWithSize(int32(varHeaderLen))
varHeaderBuf := buf.New()
defer varHeaderBuf.Release()
if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
@@ -396,21 +421,12 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
// ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream.
func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) {
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
chunkCipherLen := fixedPlainLen + AEADTagSize
headerLen := method.KeySaltLength + chunkCipherLen
// Single read call for Salt + Fixed-length response header chunk per SIP022 §3.1.4
var headerBuf [128]byte
headerSlice := headerBuf[:headerLen]
n, err := r.Read(headerSlice)
if err != nil || n < headerLen {
return nil, errors.New("failed to read complete server response header")
var serverSalt [32]byte
serverSaltSlice := serverSalt[:method.KeySaltLength]
if _, err := io.ReadFull(r, serverSaltSlice); err != nil {
return nil, err
}
serverSaltSlice := headerSlice[:method.KeySaltLength]
chunkSlice := headerSlice[method.KeySaltLength:headerLen]
sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
if err != nil {
@@ -419,6 +435,14 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
reader := NewStreamReader(r, aead)
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
chunkCipherLen := fixedPlainLen + AEADTagSize
var chunkBuf [64]byte
chunkSlice := chunkBuf[:chunkCipherLen]
if _, err := io.ReadFull(r, chunkSlice); err != nil {
return nil, err
}
decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil)
if err != nil {
return nil, errors.New("failed to decrypt server response header").Base(err)
@@ -460,190 +484,46 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
return reader, nil
}
// ServerStreamWriter lazily sends the response header along with the first payload chunk per SIP022 §3.1.2 & §3.1.4.
type ServerStreamWriter struct {
mu sync.Mutex
w io.Writer
method *CipherMethod
psk []byte
clientSalt []byte
streamWriter *StreamWriter
}
func NewServerStreamWriter(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte) *ServerStreamWriter {
return &ServerStreamWriter{
w: w,
method: method,
psk: psk,
clientSalt: clientSalt,
}
}
func (s *ServerStreamWriter) sendHeaderWithFirstPayload(payload []byte) (*StreamWriter, error) {
// WriteTCPResponse writes the server handshake response and returns a body writer for server stream.
func WriteTCPResponse(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte, initialPayload []byte) (buf.Writer, error) {
var serverSalt [32]byte
serverSaltSlice := serverSalt[:s.method.KeySaltLength]
serverSaltSlice := serverSalt[:method.KeySaltLength]
if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil {
return nil, err
}
respKey := DeriveSessionSubKey(s.psk, serverSaltSlice, s.method.KeySaltLength)
respAead, err := s.method.NewAEAD(respKey)
respKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
respAead, err := method.NewAEAD(respKey)
if err != nil {
return nil, err
}
sw := NewStreamWriter(s.w, respAead)
writer := NewStreamWriter(w, respAead)
totalHeaderLen := int32(s.method.KeySaltLength + 1 + 8 + s.method.KeySaltLength + 2 + AEADTagSize + len(payload) + AEADTagSize)
outBuf := buf.NewWithSize(totalHeaderLen)
defer outBuf.Release()
respBuf := buf.New()
defer respBuf.Release()
outBuf.Write(serverSaltSlice)
respBuf.Write(serverSaltSlice)
var fixedRespPlain [1 + 8 + 32 + 2]byte
fixedRespSlice := fixedRespPlain[:1+8+s.method.KeySaltLength+2]
fixedRespSlice := fixedRespPlain[:1+8+method.KeySaltLength+2]
fixedRespSlice[0] = HeaderTypeServer
binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix()))
copy(fixedRespSlice[9:9+s.method.KeySaltLength], s.clientSalt)
binary.BigEndian.PutUint16(fixedRespSlice[9+s.method.KeySaltLength:11+s.method.KeySaltLength], uint16(len(payload)))
copy(fixedRespSlice[9:9+method.KeySaltLength], clientSalt)
binary.BigEndian.PutUint16(fixedRespSlice[9+method.KeySaltLength:11+method.KeySaltLength], uint16(len(initialPayload)))
fixedRespChunk := sw.cipher.Seal(nil, sw.nonce[:], fixedRespSlice, nil)
IncreaseNonce(sw.nonce[:])
outBuf.Write(fixedRespChunk)
fixedRespChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedRespSlice, nil)
IncreaseNonce(writer.nonce[:])
respBuf.Write(fixedRespChunk)
if len(payload) > 0 {
payloadChunk := sw.cipher.Seal(nil, sw.nonce[:], payload, nil)
IncreaseNonce(sw.nonce[:])
outBuf.Write(payloadChunk)
if len(initialPayload) > 0 {
initialChunk := writer.cipher.Seal(nil, writer.nonce[:], initialPayload, nil)
IncreaseNonce(writer.nonce[:])
respBuf.Write(initialChunk)
}
if _, err := s.w.Write(outBuf.Bytes()); err != nil {
if _, err := w.Write(respBuf.Bytes()); err != nil {
return nil, err
}
return sw, nil
}
func (s *ServerStreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
if mb.IsEmpty() {
return nil
}
if s.streamWriter == nil {
s.mu.Lock()
if s.streamWriter == nil {
firstBuf := mb[0]
firstBytes := firstBuf.Bytes()
chunkSize := len(firstBytes)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
firstPayload := firstBytes[:chunkSize]
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
if err != nil {
s.mu.Unlock()
buf.ReleaseMulti(mb)
return err
}
s.streamWriter = sw
firstBuf.Advance(int32(chunkSize))
if firstBuf.IsEmpty() {
firstBuf.Release()
mb = mb[1:]
}
}
s.mu.Unlock()
if len(mb) == 0 {
return nil
}
}
return s.streamWriter.WriteMultiBuffer(mb)
}
func (s *ServerStreamWriter) Write(p []byte) (int, error) {
n := len(p)
if s.streamWriter == nil {
s.mu.Lock()
if s.streamWriter == nil {
chunkSize := len(p)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
firstPayload := p[:chunkSize]
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
if err != nil {
s.mu.Unlock()
return 0, err
}
s.streamWriter = sw
p = p[chunkSize:]
}
s.mu.Unlock()
if len(p) == 0 {
return n, nil
}
}
_, err := s.streamWriter.Write(p)
return n, err
}
func (s *ServerStreamWriter) Close() error {
if s.streamWriter == nil {
s.mu.Lock()
defer s.mu.Unlock()
if s.streamWriter == nil {
sw, err := s.sendHeaderWithFirstPayload(nil)
if err != nil {
return err
}
s.streamWriter = sw
}
}
return nil
}
// InitServerStream decrypts the client request header, verifies the timestamp and replay filter,
// and returns a StreamReader for subsequent stream chunks.
func InitServerStream(conn net.Conn, method *CipherMethod, psk, saltSlice []byte, salt [32]byte, fixedChunk []byte, saltFilter *antireplay.ReplayFilter[[32]byte]) (*StreamReader, *ClientRequestHeader, error) {
sessionKey := DeriveSessionSubKey(psk, saltSlice, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
if err != nil {
return nil, nil, err
}
reader := NewStreamReader(conn, aead)
reqHeader, err := ReadClientRequestHeaderWithFixed(reader, fixedChunk)
if err != nil {
return nil, nil, err
}
_ = conn.SetReadDeadline(time.Time{})
if !saltFilter.Check(salt) {
return nil, nil, ErrSaltNotUnique
}
return reader, reqHeader, nil
}
func TransportTCP(ctx context.Context, sessionPolicy policy.Session, reader buf.Reader, writer buf.Writer, link *transport.Link) error {
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
if c, ok := writer.(io.Closer); ok {
defer c.Close()
}
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
return writer, nil
}
+28 -13
View File
@@ -15,27 +15,28 @@ Plainly enabling it in the config probably will result nothing, or lock your rou
## DETAILS
By default, enabling the feature will only bring the tun interface up. \
When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS uses the first IPv4 prefix from `gateway` to configure the utun point-to-point address. \
When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS and FreeBSD use the first IPv4 prefix from `gateway` for the point-to-point address. \
Without `gateway`, the systems differ: Xray assigns no address on Linux, Windows gives the interface link-local addresses itself (an IPv6 one at once, an IPv4 one from `169.254.0.0/16` after a few seconds), and macOS and FreeBSD use `169.254.10.1/30`. \
Windows, Linux and macOS can also apply system routes from `autoSystemRoutingTable`.
macOS does not configure system DNS from the `dns` field, and neither does Linux by default; system DNS remains managed by the OS or distribution-specific network services. \
For more advanced routing policies or rules, OS level configuration can still manage the named interface (e.g. xray0) when it appears.
This keeps complex system level routing and rules in a single place of responsibility - the OS itself. \
Examples of how to achieve this on a simple Linux system (Ubuntu with systemd-networkd) can be found at the end of this README.
### SYSTEM DNS ON LINUX (`autoSystemDNS`)
### SYSTEM DNS ON LINUX (`autoSystemDnsToGateway`)
On Linux, setting `autoSystemDNS` to `true` lets the inbound point the system resolver at the tun interface, so name lookups resolve through Xray instead of going out over the physical link. It is off by default, and it is Linux-only.
On Linux, setting `autoSystemDnsToGateway` to `true` lets the inbound point the system resolver at the tun interface, so name lookups resolve through Xray instead of going out over the physical link. It is off by default, and it is Linux-only.
It uses `resolvectl`, which means it applies only when all of these hold:
It uses `resolvectl`, which means it only works when all of these hold. Where Xray can tell that one does not, it does not start:
- the system runs systemd and `resolvectl` is on `PATH`
- `systemd-resolved` is enabled and actually managing DNS (installed but not running has no effect)
- `systemd-resolved` is enabled and actually managing DNS (installed but not running is not enough)
- systemd-resolved is version 240 or newer, where `default-route` exists
- no `dns` upstream resolves through the system resolver, directly or through its own bootstrap (see below)
The address handed over is the first IPv4 `gateway` incremented by one (e.g. `192.168.100.1/30` -> `192.168.100.2`). It is not taken from `dns`: handing `1.1.1.1` to `resolvectl dns` would make systemd-resolved query that server directly over the physical link, which is the leak this option exists to close.
The address handed over is the first IPv4 `gateway`, or without one the first IPv6 `gateway`, incremented by one (e.g. `192.168.100.1/30` -> `192.168.100.2`, `fc00::1/64` -> `fc00::2`). Without any `gateway`, the config is rejected. It is not taken from `dns`: handing `1.1.1.1` to `resolvectl dns` would make systemd-resolved query that server directly over the physical link, which is the leak this option exists to close.
Because that address has to actually answer, the takeover is checked before it happens. A query from the interface address to that address is routed through the configured rules, and host-wide DNS is only changed when the result is a DNS-capable outbound. Otherwise the option does nothing and DNS is left to the OS. In practice this means you also need a routing rule sending the interface's port 53 to a `dns` outbound, for example:
Because that address has to actually answer, the takeover is checked before it happens. A query from the interface address to that address is routed through the configured rules, and host-wide DNS is only changed when the result is a DNS-capable outbound. Otherwise DNS is left alone and Xray does not start. In practice this means you also need a routing rule sending the interface's port 53 to a `dns` outbound, for example:
```json
"routing": {
@@ -49,19 +50,19 @@ The check is a preflight, not a proof for arbitrary rules. It sends its query fr
It is also a check for the dependencies it knows about, not a proof that no indirect one exists. A hostname-based upstream that bootstraps through system DNS is the case in point: `https+local://dns.google/dns-query` resolves its own hostname with `DialSystem`, so once the takeover is in place that bootstrap goes `resolved -> TUN -> DNS outbound -> bootstrap -> resolved` and the query times out. The preflight does not see it, because the dependency sits in the upstream's bootstrap rather than in the clients it inspects. Upstream resolution, bootstrap included, therefore has to stay independent of the resolver path being redirected; configuring the address instead of the hostname, or resolving the hostname beforehand, avoids it.
The upstream requirement in the list above matters as much as the routing rule. With no name servers configured, Core resolves through a client that forwards to the system resolver; pointing the system resolver at the TUN would then close a loop through the DNS outbound, `resolved -> TUN -> DNS outbound -> system resolver -> resolved`, and resolution stops. The takeover is refused in that case.
The upstream requirement in the list above matters as much as the routing rule. With no name servers configured, Core resolves through a client that forwards to the system resolver; pointing the system resolver at the TUN would then close a loop through the DNS outbound, `resolved -> TUN -> DNS outbound -> system resolver -> resolved`, and resolution stops. The takeover is refused in that case, and Xray does not start.
The same applies to a name server pointed at `localhost`, and to a `dns` section that is present but lists no name servers. One such upstream is enough to refuse the takeover even when independent upstreams are configured alongside it: name servers are selected per domain, so a domain-specific rule can still choose the local one, and the loop then affects whichever domains reach it. The check is deliberately broader than the loop it observed, because the alternative would be to drop a name server the user configured.
Where it does not apply, DNS is left alone and the leak described in XTLS/Xray-core#6454 remains:
Where it cannot apply, Xray does not start, rather than run with the leak described in XTLS/Xray-core#6454, so leave the option off there:
| Environment | Behaviour |
|---|---|
| systemd distribution with systemd-resolved enabled | applies |
| Alpine, Void, Devuan, OpenRC-based, OpenWrt | no `resolvectl`, skipped |
| DNS managed by dnsmasq / unbound / BIND / static `resolv.conf` | unreachable by `resolvectl`, skipped |
| Containers without a systemd-resolved daemon | skipped |
| systemd older than 240 | `default-route` unavailable, skipped |
| Alpine, Void, Devuan, OpenRC-based, OpenWrt | no `resolvectl`, does not start |
| DNS managed by dnsmasq / unbound / BIND / static `resolv.conf` | unreachable by `resolvectl`, does not start |
| Containers without a systemd-resolved daemon | does not start |
| systemd older than 240 | `default-route` unavailable, does not start |
On `Close()` the setting is reverted. It is **not** reverted if the process is killed with `SIGKILL`, since a process cannot handle that signal; run `resolvectl revert <iface>` to clean up by hand. An application that brings its own DNS endpoint is unaffected either way — this only covers the system resolver.
@@ -198,6 +199,20 @@ To make it start, wintun.dll specific for your Windows/arch must be present next
After the start network adapter with the name you chose in the config will be created in the system, and exist while Xray is running.
When `dns` is set, those servers are applied to the adapter. Windows is kept from registering the TUN's addresses in DNS, and its DNS cache is flushed when the TUN starts and stops.
With `autoSystemWfpBlockLeak`, which needs `autoSystemRoutingTable` (the config is rejected otherwise), Xray also adds Windows Filtering Platform filters that keep two kinds of traffic of every program but Xray itself from leaving outside the TUN, each chosen by a value in the list, e.g. `"autoSystemWfpBlockLeak": ["dns", "misconfigtun"]`:
- `"dns"` (needs `dns`, the config is rejected otherwise): DNS (port 53) only goes through the TUN. Windows keeps sending name queries to the DNS servers of the other interfaces as well, out through those interfaces whatever the routes say, and other programs reach a resolver on the local network (e.g. `192.168.1.1` handed out by DHCP) through its more specific LAN route instead of the TUN. On Windows 11 and Server 2022 and later, where those queries may also go over HTTPS or TLS, Windows' DNS Client service cannot connect outside the TUN at all, except for name resolution on the local network (LLMNR, mDNS). The `dns` servers therefore have to lie within `gateway` or `autoSystemRoutingTable` (a warning is logged otherwise), and DNS servers that should be reached directly belong in Xray's own `dns` settings.
- `"misconfigtun"`: an IP version without routes in `autoSystemRoutingTable`, IPv4 or IPv6, is blocked entirely, in both directions, as it would bypass the TUN. Only loopback and what Windows itself needs on the local link (DHCP, and for IPv6 neighbor and multicast listener discovery) remain allowed. An address of that version in `gateway` is not needed: without one, Windows gives the TUN link-local addresses itself, an IPv6 one at once and an IPv4 one from `169.254.0.0/16` after some seconds (until then, IPv4 routed to the TUN is unreachable), and what is routed to the TUN goes through it with those.
With the filters in place, Xray's own connections out also get past Windows Firewall's block rules (other firewalls may still block them), while connections to Xray's inbounds stay subject to them.
Names that Xray resolves through the system resolver, such as an outbound's server address given as a domain with the default `AsIs` domain strategy, would be looked up by Windows on Xray's behalf, and those queries would then go into the TUN too. While DNS is restricted this way and `autoOutboundsInterface` is in use (the default with `autoSystemRoutingTable`), Xray therefore resolves them itself, with its own queries to the DNS servers of the other interfaces. That bypasses Windows' DNS cache, and its name resolution on the local network (LLMNR, mDNS): a server address given as a domain is looked up again for every connection, and a DNS server that does not answer delays each lookup. Having Xray's own `dns` resolve it, through the outbound's `sockopt.domainStrategy`, avoids that. The `localhost` DNS server queries the same servers whenever `autoOutboundsInterface` is in use. Both skip the TUN's own DNS servers, unless another interface uses them as well: queried from Xray itself, they would lead back into it, or nowhere.
If the filters cannot be added, Xray does not start. They are removed when Xray exits. Not covered is name resolution on the local network (LLMNR, mDNS, NetBIOS), except over an IP version that is blocked.
`autoSystemWfpBlockLeak` (Windows only) is empty by default, as the filters break some setups: with `"dns"`, a local DNS resolver other programs use (e.g. on `127.0.0.1:53`), the DNS of another VPN on its own interface, virtual machines whose NAT resolves names on the host, or signing in to a captive portal; with `"misconfigtun"`, IPv4 or IPv6 on the local network while no route of that version leads to the TUN. Without the filters, DNS may leak as described above. To keep an IP version out of the TUN on purpose while still blocking DNS leaks, use only `["dns"]`.
You can give the adapter ip address manually, you can live Windows to give it autogenerated ip address (which take few seconds), it doesn't matter, the traffic going _through_ the interface will be forwarded into the app for proxying. \
Minimal configuration that will work for local machine is routing passing the traffic on-link through the interface.
You will need the interface id for that, unfortunately it is going to change with every Xray start due to implementation ambiguity between Xray and wintun driver.
+16 -6
View File
@@ -32,7 +32,8 @@ type Config struct {
AutoSystemRoutingTable []string `protobuf:"bytes,6,rep,name=auto_system_routing_table,json=autoSystemRoutingTable,proto3" json:"auto_system_routing_table,omitempty"`
AutoOutboundsInterface string `protobuf:"bytes,7,opt,name=auto_outbounds_interface,json=autoOutboundsInterface,proto3" json:"auto_outbounds_interface,omitempty"`
Desc string `protobuf:"bytes,8,opt,name=desc,proto3" json:"desc,omitempty"`
AutoSystemDns bool `protobuf:"varint,9,opt,name=auto_system_dns,json=autoSystemDns,proto3" json:"auto_system_dns,omitempty"`
AutoSystemDnsToGateway bool `protobuf:"varint,9,opt,name=auto_system_dns_to_gateway,json=autoSystemDnsToGateway,proto3" json:"auto_system_dns_to_gateway,omitempty"`
AutoSystemWfpBlockLeak []string `protobuf:"bytes,10,rep,name=auto_system_wfp_block_leak,json=autoSystemWfpBlockLeak,proto3" json:"auto_system_wfp_block_leak,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
@@ -123,18 +124,25 @@ func (x *Config) GetDesc() string {
return ""
}
func (x *Config) GetAutoSystemDns() bool {
func (x *Config) GetAutoSystemDnsToGateway() bool {
if x != nil {
return x.AutoSystemDns
return x.AutoSystemDnsToGateway
}
return false
}
func (x *Config) GetAutoSystemWfpBlockLeak() []string {
if x != nil {
return x.AutoSystemWfpBlockLeak
}
return nil
}
var File_proxy_tun_config_proto protoreflect.FileDescriptor
const file_proxy_tun_config_proto_rawDesc = "" +
"\n" +
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\xaa\x02\n" +
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\xfa\x02\n" +
"\x06Config\x12\x12\n" +
"\x04name\x18\x01 \x01(\tR\x04name\x12\x10\n" +
"\x03MTU\x18\x02 \x01(\rR\x03MTU\x12\x18\n" +
@@ -144,8 +152,10 @@ const file_proxy_tun_config_proto_rawDesc = "" +
"user_level\x18\x05 \x01(\rR\tuserLevel\x129\n" +
"\x19auto_system_routing_table\x18\x06 \x03(\tR\x16autoSystemRoutingTable\x128\n" +
"\x18auto_outbounds_interface\x18\a \x01(\tR\x16autoOutboundsInterface\x12\x12\n" +
"\x04desc\x18\b \x01(\tR\x04desc\x12&\n" +
"\x0fauto_system_dns\x18\t \x01(\bR\rautoSystemDnsBL\n" +
"\x04desc\x18\b \x01(\tR\x04desc\x12:\n" +
"\x1aauto_system_dns_to_gateway\x18\t \x01(\bR\x16autoSystemDnsToGateway\x12:\n" +
"\x1aauto_system_wfp_block_leak\x18\n" +
" \x03(\tR\x16autoSystemWfpBlockLeakBL\n" +
"\x12com.xray.proxy.tunP\x01Z#github.com/xtls/xray-core/proxy/tun\xaa\x02\x0eXray.Proxy.Tunb\x06proto3"
var (
+2 -1
View File
@@ -15,5 +15,6 @@ message Config {
repeated string auto_system_routing_table = 6;
string auto_outbounds_interface = 7;
string desc = 8;
bool auto_system_dns = 9;
bool auto_system_dns_to_gateway = 9;
repeated string auto_system_wfp_block_leak = 10;
}
+4 -2
View File
@@ -166,12 +166,14 @@ func (t *Handler) Start() error {
}
// Platform-specific system DNS takeover, where the platform implements it.
// Non-fatal: a failure leaves DNS management with the OS.
// Rather no TUN than one that the system DNS bypasses.
if c, ok := tunInterface.(interface {
ConfigureSystemDNS(context.Context, string) error
}); ok {
if err := c.ConfigureSystemDNS(t.ctx, t.tag); err != nil {
errors.LogInfoInner(t.ctx, err, "[tun] system DNS not configured")
_ = tunStack.Close()
_ = tunInterface.Close()
return errors.New("unable to set the system DNS (remove autoSystemDnsToGateway to run without)").Base(err)
}
}
+21 -15
View File
@@ -53,23 +53,29 @@ var resolvectlRunner = func(name string, args ...string) ([]byte, error) {
}
// systemDNSAddrs derives the addresses used for the system DNS takeover from the
// first IPv4 gateway: the gateway address itself is what a query from this
// interface appears to come from, and the next address is what the resolver is
// pointed at. The latter belongs to the TUN and is answered inside Xray;
// handing the configured public resolvers to resolvectl instead would leave the
// system querying them directly over the physical link, defeating the point of
// the TUN.
// first IPv4 gateway, or without one, the first IPv6 gateway: the gateway
// address itself is what a query from this interface appears to come from, and
// the next address is what the resolver is pointed at. The latter belongs to
// the TUN and is answered inside Xray; handing the configured public resolvers
// to resolvectl instead would leave the system querying them directly over the
// physical link, defeating the point of the TUN.
func systemDNSAddrs(gateway []string) (source, dns netip.Addr, ok bool) {
var first6 netip.Addr
for _, address := range gateway {
prefix, err := netip.ParsePrefix(address)
if err != nil {
continue
}
addr := prefix.Addr()
if !addr.Is4() {
continue
if addr.Is4() {
return addr, addr.Next(), true
}
return addr, addr.Next(), true
if !first6.IsValid() {
first6 = addr
}
}
if first6.IsValid() {
return first6, first6.Next(), true
}
return netip.Addr{}, netip.Addr{}, false
}
@@ -115,11 +121,11 @@ const probeSourcePort = 49152
// Overridable for tests.
var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address string) error {
ip, err := netip.ParseAddr(address)
if err != nil || !ip.Is4() {
if err != nil {
return errors.New("invalid DNS address ", address).Base(err)
}
src, err := netip.ParseAddr(source)
if err != nil || !src.Is4() {
if err != nil || src.Is4() != ip.Is4() {
return errors.New("invalid source address ", source).Base(err)
}
@@ -182,10 +188,10 @@ var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address str
//
// It acts only when the config opts in, and it verifies the data path first:
// unless a query to the advertised address would actually be handled, host-wide
// resolution is left to the OS, which is the documented default. Errors are
// returned to the caller, which treats them as non-fatal.
// resolution is left to the OS and an error returned. The caller does not start
// the TUN on an error, as the system DNS would bypass it.
func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) error {
if !t.options.AutoSystemDns {
if !t.options.AutoSystemDnsToGateway {
return nil
}
if t.systemDNSSet {
@@ -202,7 +208,7 @@ func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) er
source, address, ok := systemDNSAddrs(t.options.Gateway)
if !ok {
return errors.New("no IPv4 gateway, cannot derive a system DNS address")
return errors.New("no gateway, cannot derive a system DNS address")
}
iface := t.ifaceName()
+12
View File
@@ -191,3 +191,15 @@ func TestVerifyDNSRoutingDecisions(t *testing.T) {
})
}
}
// Without an IPv4 gateway, the takeover uses the first IPv6 one, and the probe
// carries IPv6 addresses.
func TestVerifyDNSRoutingIPv6(t *testing.T) {
ctx := newRouteTestContext(t, true, udpNameServer([]byte{9, 9, 9, 9}), []*router.RoutingRule{port53Rule()})
if err := verifyDNSRouting(ctx, routeTestInboundTag, "fc00::1", "fc00::2"); err != nil {
t.Fatalf("expected the takeover to be accepted, got: %v", err)
}
if err := verifyDNSRouting(ctx, routeTestInboundTag, routeTestSource, "fc00::2"); err == nil {
t.Fatal("expected mixed IPv4 and IPv6 addresses to be refused")
}
}
+17 -8
View File
@@ -58,9 +58,9 @@ func recorder(t *testing.T, failOn string) *[][]string {
func optedInTun() *LinuxTun {
return &LinuxTun{
options: &Config{
Name: "xray_tun",
Gateway: []string{"192.168.100.1/30"},
AutoSystemDns: true,
Name: "xray_tun",
Gateway: []string{"192.168.100.1/30"},
AutoSystemDnsToGateway: true,
},
tunLink: testLink("xray_tun"),
}
@@ -79,7 +79,7 @@ func TestConfigureSystemDNSDisabledByDefault(t *testing.T) {
calls := recorder(t, "")
t1 := optedInTun()
t1.options.AutoSystemDns = false
t1.options.AutoSystemDnsToGateway = false
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
t.Fatalf("unexpected error: %v", err)
@@ -103,7 +103,7 @@ func TestConfigureSystemDNSNoGateway(t *testing.T) {
t1.options.Gateway = nil
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
t.Fatal("expected an error when no IPv4 gateway is configured")
t.Fatal("expected an error when no gateway is configured")
}
if len(*probes) != 0 {
t.Errorf("routing probe must not run without a gateway, got %d calls", len(*probes))
@@ -351,9 +351,18 @@ func TestSystemDNSAddrs(t *testing.T) {
wantOK: false,
},
{
name: "ipv6 only",
gateway: []string{"fc00::1/64"},
wantOK: false,
name: "ipv6 only",
gateway: []string{"fc00::1/64"},
wantSource: "fc00::1",
wantDNS: "fc00::2",
wantOK: true,
},
{
name: "first ipv6 without ipv4",
gateway: []string{"fc00::1/64", "fd00::1/64"},
wantSource: "fc00::1",
wantDNS: "fc00::2",
wantOK: true,
},
}
+229 -2
View File
@@ -3,17 +3,25 @@
package tun
import (
"bytes"
"context"
"crypto/md5"
"encoding/binary"
go_errors "errors"
"net"
"net/netip"
"os/exec"
"path/filepath"
"slices"
"strconv"
"strings"
"sync"
"syscall"
"time"
"unsafe"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet"
"golang.org/x/sys/windows"
"golang.zx2c4.com/wintun"
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
@@ -38,6 +46,10 @@ type WindowsTun struct {
luid winipcfg.LUID
cbr winipcfg.ChangeCallback
cbi winipcfg.ChangeCallback
wfp windows.Handle
resolver *savedResolver
skipStop chan struct{}
skipDone chan struct{}
closed bool
}
@@ -197,19 +209,105 @@ startOver:
}
}
// Windows lists the TUN's DNS servers among the system's ones, which Go's
// resolver queries for Xray's own lookups past the TUN, where they lead
// nowhere or back into Xray. Not skipped are those another interface uses
// as well, as that could leave no server at all. As those can change at
// any time, they are looked at again as often as Go rereads its servers.
if len(dns) > 0 {
skipped, err := tunOnlyDNS(t.luid, dns)
if err != nil {
skipped = dns
}
internet.SkipDNSServers(skipped)
t.skipStop, t.skipDone = make(chan struct{}), make(chan struct{})
go func() {
defer close(t.skipDone)
ticker := time.NewTicker(5 * time.Second)
defer ticker.Stop()
for {
select {
case <-ticker.C:
if skipped, err := tunOnlyDNS(t.luid, dns); err == nil {
internet.SkipDNSServers(skipped)
}
case <-t.skipStop:
return
}
}
}()
}
// Keep Windows from registering the TUN's addresses, and the host name
// with them, through dynamic DNS updates. Best effort.
if address4 || address6 {
if err := disableDNSRegistration(t.luid, dns); err != nil {
errors.LogDebugInner(context.Background(), err, "[tun] unable to disable DNS registration")
}
}
// With autoSystemWfpBlockLeak, once the system routes lead to the TUN,
// keep DNS ("dns", if dns is set), and an IP version no route of which
// leads to the TUN ("misconfigtun"), from leaving through the other
// interfaces. Addresses do not matter: without one of a version in
// gateway, Windows gives the TUN a link-local one.
leaks := t.options.AutoSystemWfpBlockLeak
blockDNS := slices.Contains(leaks, "dns") && len(dns) > 0
blockIPv4 := slices.Contains(leaks, "misconfigtun") && !route4
blockIPv6 := slices.Contains(leaks, "misconfigtun") && !route6
if (route4 || route6) && (blockDNS || blockIPv4 || blockIPv6) {
if t.wfp, err = blockLeaks(t.luid, blockDNS, blockIPv4, blockIPv6); err != nil {
var blocked []string
for _, b := range []struct {
on bool
what string
}{{blockDNS, "DNS"}, {blockIPv4, "IPv4"}, {blockIPv6, "IPv6"}} {
if b.on {
blocked = append(blocked, b.what)
}
}
// Rather no TUN than a leaking one.
return errors.New("unable to block ", strings.Join(blocked, " and "), " outside the TUN (remove autoSystemWfpBlockLeak to run without)").Base(err)
}
errors.LogInfo(context.Background(), "[tun] outside the TUN, blocked DNS: ", blockDNS, ", blocked IPv4: ", blockIPv4, ", blocked IPv6: ", blockIPv6)
if blockDNS {
covered := slices.Clone(addresses)
for _, route := range routesData {
covered = append(covered, route.Destination)
}
for _, server := range dnsOutsideTUN(dns, covered) {
errors.LogWarning(context.Background(), "[tun] DNS server ", server, " is in neither gateway nor autoSystemRoutingTable, so queries to it cannot go through the TUN and are blocked")
}
// With updater, the dialer controllers bind Xray's own sockets
// to the physical interface.
if updater != nil {
t.resolver = resolveOnOwn()
}
}
}
if len(dns) > 0 || route4 || route6 {
if err := flushDNSCache(); err != nil {
errors.LogInfoInner(context.Background(), err, "[tun] unable to flush DNS cache")
}
}
if updater != nil {
t.cbr, err = winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) {
// Only a registered callback goes into the fields: a nil pointer in
// them would not compare equal to nil in Close.
cbr, err := winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) {
updater.Update()
})
if err != nil {
return err
}
t.cbi, err = winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
t.cbr = cbr
cbi, err := winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
updater.Update()
})
if err != nil {
return err
}
t.cbi = cbi
}
return nil
}
@@ -236,6 +334,20 @@ func (t *WindowsTun) Close() error {
t.luid.FlushIPAddresses(windows.AF_INET6)
t.luid.FlushDNS(windows.AF_INET6)
}
if t.wfp != 0 {
closeWFPEngine(t.wfp)
}
if t.resolver != nil {
t.resolver.restore()
}
if t.skipStop != nil {
close(t.skipStop)
<-t.skipDone
}
internet.SkipDNSServers(nil)
if len(t.options.DNS) > 0 || len(t.options.AutoSystemRoutingTable) > 0 {
flushDNSCache()
}
if t.session != (wintun.Session{}) {
t.session.End()
}
@@ -245,6 +357,121 @@ func (t *WindowsTun) Close() error {
return nil
}
type savedResolver struct {
preferGo bool
dial func(ctx context.Context, network, address string) (net.Conn, error)
}
// resolveOnOwn has Go resolve the names Xray would otherwise ask Windows for,
// on Xray's own sockets, which the dialer controllers bind to the physical
// interface, and skipping the TUN's DNS servers, as localdns does. Windows'
// resolver runs in the DNS Client service, whose queries the DNS filter lets
// through the TUN only, so Xray's own lookups, like of an outbound's server
// domain, would go into Xray again and could end up waiting on themselves.
//
// It changes net.DefaultResolver for the whole process, which covers every
// lookup that would reach Windows' resolver; restore undoes it.
func resolveOnOwn() *savedResolver {
saved := &savedResolver{net.DefaultResolver.PreferGo, net.DefaultResolver.Dial}
dialer := &net.Dialer{Control: func(network, address string, c syscall.RawConn) error {
for _, ctl := range internet.Controllers {
if err := ctl(network, address, c); err != nil {
return err
}
}
return nil
}}
// Go's resolver moves on to the next server right away when a dial fails.
net.DefaultResolver.Dial = func(ctx context.Context, network, address string) (net.Conn, error) {
if internet.IsSkippedDNSServer(address) {
return nil, errors.New("skipped DNS server ", address)
}
return dialer.DialContext(ctx, network, address)
}
net.DefaultResolver.PreferGo = true
return saved
}
func (s *savedResolver) restore() {
net.DefaultResolver.PreferGo = s.preferGo
net.DefaultResolver.Dial = s.dial
}
// tunOnlyDNS returns those of servers, the TUN's DNS servers, that Go's
// resolver does not also get from another interface: one that is up and has
// a gateway, as it reads them.
func tunOnlyDNS(tun winipcfg.LUID, servers []netip.Addr) ([]netip.Addr, error) {
adapters, err := winipcfg.GetAdaptersAddresses(windows.AF_UNSPEC, winipcfg.GAAFlagIncludeGateways)
if err != nil {
return nil, err
}
var others []netip.Addr
for _, adapter := range adapters {
if adapter.LUID == tun || adapter.OperStatus != winipcfg.IfOperStatusUp || adapter.FirstGatewayAddress == nil {
continue
}
for server := adapter.FirstDNSServerAddress; server != nil; server = server.Next {
if addr, ok := netip.AddrFromSlice(server.Address.IP()); ok {
others = append(others, addr.Unmap())
}
}
}
return slices.DeleteFunc(slices.Clone(servers), func(server netip.Addr) bool {
return slices.Contains(others, server.Unmap())
}), nil
}
// disableDNSRegistration turns off the dynamic DNS registration of the
// interface's addresses. dns are its DNS servers.
func disableDNSRegistration(luid winipcfg.LUID, dns []netip.Addr) error {
guid, err := luid.GUID()
if err != nil {
return err
}
err = winipcfg.SetInterfaceDnsSettings(*guid, &winipcfg.DnsInterfaceSettings{
Version: winipcfg.DnsInterfaceSettingsVersion1,
Flags: winipcfg.DnsInterfaceSettingsFlagRegistrationEnabled,
})
if err == nil || !go_errors.Is(err, windows.ERROR_PROC_NOT_FOUND) {
return err
}
return disableDNSRegistrationByNetsh(luid, dns)
}
// disableDNSRegistrationByNetsh does it for Windows before 10 1809, which
// lacks SetInterfaceDnsSettings. The setting is the interface's, not the
// address family's, but netsh only applies it along with a DNS server, which
// replaces the IPv4 ones, so they are set again afterwards.
func disableDNSRegistrationByNetsh(luid winipcfg.LUID, dns []netip.Addr) error {
row, err := luid.Interface()
if err != nil {
return err
}
server := "127.0.0.1" // any will do when there is no IPv4 one
if i := slices.IndexFunc(dns, netip.Addr.Is4); i >= 0 {
server = dns[i].String()
}
err = runNetsh("interface", "ipv4", "set", "dnsservers", "name="+strconv.FormatUint(uint64(row.InterfaceIndex), 10), "source=static", "address="+server, "register=none", "validate=no")
return errors.Combine(err, luid.SetDNS(windows.AF_INET, dns, nil))
}
// runNetsh runs netsh.exe from the system directory. netsh reports some
// failures, like a syntax error, only in its output, even with exit code 0,
// so any output counts as a failure.
func runNetsh(args ...string) error {
system32, err := windows.GetSystemDirectory()
if err != nil {
return err
}
cmd := exec.Command(filepath.Join(system32, "netsh.exe"), args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
output, err := cmd.CombinedOutput()
if output = bytes.TrimSpace(output); err != nil || len(output) > 0 {
return errors.New("netsh ", strings.Join(args, " "), ": ", string(output)).Base(err)
}
return nil
}
func (t *WindowsTun) Name() (string, error) {
row, err := t.luid.Interface()
if err != nil {
+471
View File
@@ -0,0 +1,471 @@
//go:build windows
package tun
import (
"net/netip"
"os"
"runtime"
"slices"
"unsafe"
"github.com/xtls/xray-core/common/errors"
"golang.org/x/sys/windows"
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
)
var (
modfwpuclnt = windows.NewLazySystemDLL("fwpuclnt.dll")
moddnsapi = windows.NewLazySystemDLL("dnsapi.dll")
procFwpmEngineOpen0 = modfwpuclnt.NewProc("FwpmEngineOpen0")
procFwpmEngineClose0 = modfwpuclnt.NewProc("FwpmEngineClose0")
procFwpmTransactionBegin0 = modfwpuclnt.NewProc("FwpmTransactionBegin0")
procFwpmTransactionCommit0 = modfwpuclnt.NewProc("FwpmTransactionCommit0")
procFwpmTransactionAbort0 = modfwpuclnt.NewProc("FwpmTransactionAbort0")
procFwpmSubLayerAdd0 = modfwpuclnt.NewProc("FwpmSubLayerAdd0")
procFwpmFilterAdd0 = modfwpuclnt.NewProc("FwpmFilterAdd0")
procFwpmGetAppIdFromFileName0 = modfwpuclnt.NewProc("FwpmGetAppIdFromFileName0")
procFwpmFreeMemory0 = modfwpuclnt.NewProc("FwpmFreeMemory0")
procDnsFlushResolverCache = moddnsapi.NewProc("DnsFlushResolverCache")
)
// fwptypes.h and fwpmtypes.h
const (
rpcCAuthnWinNT = 10 // RPC_C_AUTHN_WINNT
fwpmSessionFlagDynamic = 1 // FWPM_SESSION_FLAG_DYNAMIC
fwpmFilterFlagClearActionRight = 8 // FWPM_FILTER_FLAG_CLEAR_ACTION_RIGHT
fwpUint8 = 1 // FWP_UINT8
fwpUint16 = 2 // FWP_UINT16
fwpUint32 = 3 // FWP_UINT32
fwpUint64 = 4 // FWP_UINT64
fwpByteArray16Type = 11 // FWP_BYTE_ARRAY16_TYPE
fwpByteBlobType = 12 // FWP_BYTE_BLOB_TYPE
fwpSecurityDescriptorType = 14 // FWP_SECURITY_DESCRIPTOR_TYPE
fwpMatchEqual = 0 // FWP_MATCH_EQUAL
fwpMatchFlagsAllSet = 6 // FWP_MATCH_FLAGS_ALL_SET
fwpConditionFlagIsLoopback = 1 // FWP_CONDITION_FLAG_IS_LOOPBACK
fwpActionBlock = 0x1001 // FWP_ACTION_BLOCK
fwpActionPermit = 0x1002 // FWP_ACTION_PERMIT
)
// fwpmu.h
var (
fwpmLayerALEAuthConnectV4 = windows.GUID{Data1: 0xc38d57d1, Data2: 0x05a7, Data3: 0x4c33, Data4: [8]byte{0x90, 0x4f, 0x7f, 0xbc, 0xee, 0xe6, 0x0e, 0x82}}
fwpmLayerALEAuthConnectV6 = windows.GUID{Data1: 0x4a72393b, Data2: 0x319f, Data3: 0x44bc, Data4: [8]byte{0x84, 0xc3, 0xba, 0x54, 0xdc, 0xb3, 0xb6, 0xb4}}
fwpmLayerALEAuthRecvAcceptV4 = windows.GUID{Data1: 0xe1cd9fe7, Data2: 0xf4b5, Data3: 0x4273, Data4: [8]byte{0x96, 0xc0, 0x59, 0x2e, 0x48, 0x7b, 0x86, 0x50}}
fwpmLayerALEAuthRecvAcceptV6 = windows.GUID{Data1: 0xa3b42c97, Data2: 0x9f04, Data3: 0x4672, Data4: [8]byte{0xb8, 0x7e, 0xce, 0xe9, 0xc4, 0x83, 0x25, 0x7f}}
fwpmConditionFlags = windows.GUID{Data1: 0x632ce23b, Data2: 0x5167, Data3: 0x435c, Data4: [8]byte{0x86, 0xd7, 0xe9, 0x03, 0x68, 0x4a, 0xa8, 0x0c}}
fwpmConditionIPArrivalInterface = windows.GUID{Data1: 0x618a9b6d, Data2: 0x386b, Data3: 0x4136, Data4: [8]byte{0xad, 0x6e, 0xb5, 0x15, 0x87, 0xcf, 0xb1, 0xcd}}
fwpmConditionIPLocalInterface = windows.GUID{Data1: 0x4cd62a49, Data2: 0x59c3, Data3: 0x4969, Data4: [8]byte{0xb7, 0xf3, 0xbd, 0xa5, 0xd3, 0x28, 0x90, 0xa4}}
fwpmConditionIPLocalPort = windows.GUID{Data1: 0x0c1ba1af, Data2: 0x5765, Data3: 0x453f, Data4: [8]byte{0xaf, 0x22, 0xa8, 0xf7, 0x91, 0xac, 0x77, 0x5b}} // also FWPM_CONDITION_ICMP_TYPE
fwpmConditionIPNexthopInterface = windows.GUID{Data1: 0x93ae8f5b, Data2: 0x7f6f, Data3: 0x4719, Data4: [8]byte{0x98, 0xc8, 0x14, 0xe9, 0x74, 0x29, 0xef, 0x04}}
fwpmConditionIPProtocol = windows.GUID{Data1: 0x3971ef2b, Data2: 0x623e, Data3: 0x4f9a, Data4: [8]byte{0x8c, 0xb1, 0x6e, 0x79, 0xb8, 0x06, 0xb9, 0xa7}}
fwpmConditionIPRemoteAddress = windows.GUID{Data1: 0xb235ae9a, Data2: 0x1d64, Data3: 0x49b8, Data4: [8]byte{0xa4, 0x4c, 0x5f, 0xf3, 0xd9, 0x09, 0x50, 0x45}}
fwpmConditionIPRemotePort = windows.GUID{Data1: 0xc35a604d, Data2: 0xd22b, Data3: 0x4e1a, Data4: [8]byte{0x91, 0xb4, 0x68, 0xf6, 0x74, 0xee, 0x67, 0x4b}} // also FWPM_CONDITION_ICMP_CODE
fwpmConditionALEAppID = windows.GUID{Data1: 0xd78e1e87, Data2: 0x8644, Data3: 0x4ea5, Data4: [8]byte{0x94, 0x37, 0xd8, 0x09, 0xec, 0xef, 0xc9, 0x71}}
fwpmConditionALEUserID = windows.GUID{Data1: 0xaf043a0a, Data2: 0xb34d, Data3: 0x4f86, Data4: [8]byte{0x97, 0x9c, 0xc9, 0x03, 0x71, 0xaf, 0x6e, 0x66}}
)
// dnsClientSID is the SID of Windows' DNS Client service, NT SERVICE\Dnscache.
// Service SIDs derive from the service name, so it is the same everywhere (sc
// showsid dnscache).
const dnsClientSID = "S-1-5-80-859482183-879914841-863379149-1145462774-2388618682"
// ff02::1:2, where DHCPv6 clients send to. A package-level variable never
// moves, so conditions may refer to it through uintptr.
var ipv6AllDHCPv6Servers = [16]byte{0xff, 0x02, 13: 0x01, 15: 0x02}
type fwpByteBlob struct {
size uint32
data *byte
}
// fwpValue0 is FWP_VALUE0 as well as FWP_CONDITION_VALUE0. Their union holds
// a scalar of at most 32 bits, or a pointer for the larger types.
type fwpValue0 struct {
typ uint32
value uintptr
}
type fwpmDisplayData0 struct {
name *uint16
description *uint16
}
type fwpmSession0 struct {
sessionKey windows.GUID
displayData fwpmDisplayData0
flags uint32
txnWaitTimeoutInMSec uint32
processID uint32
sid *windows.SID
username *uint16
kernelMode int32
}
type fwpmSublayer0 struct {
subLayerKey windows.GUID
displayData fwpmDisplayData0
flags uint32
providerKey *windows.GUID
providerData fwpByteBlob
weight uint16
}
type fwpmFilterCondition0 struct {
fieldKey windows.GUID
matchType uint32
conditionValue fwpValue0
}
type fwpmAction0 struct {
typ uint32
filterType windows.GUID
}
type fwpmFilter0 struct {
filterKey windows.GUID
displayData fwpmDisplayData0
flags uint32
providerKey *windows.GUID
providerData fwpByteBlob
layerKey windows.GUID
subLayerKey windows.GUID
weight fwpValue0
numFilterConditions uint32
filterCondition *fwpmFilterCondition0
action fwpmAction0
_ uint32 // C aligns the following union to 8 bytes, as it holds a UINT64
providerContextKey windows.GUID
reserved *windows.GUID
_ [8 - unsafe.Sizeof(uintptr(0))]byte // and filterId as well, also on 32-bit
filterID uint64
effectiveWeight fwpValue0
}
// fwpmResult converts the DWORD status the Fwpm functions return.
func fwpmResult(r1, _ uintptr, _ error) error {
if r1 != 0 {
return windows.Errno(r1)
}
return nil
}
func utf16Ptr(s string) *uint16 {
p, _ := windows.UTF16PtrFromString(s)
return p
}
func condition(field *windows.GUID, typ uint32, value uintptr) fwpmFilterCondition0 {
return fwpmFilterCondition0{
fieldKey: *field,
matchType: fwpMatchEqual,
conditionValue: fwpValue0{typ: typ, value: value},
}
}
// blockLeaks keeps traffic from leaving through interfaces other than tun,
// for every program but Xray itself, whose outbounds (DNS included) use the
// other interfaces on purpose:
//
// - dns: DNS (port 53) may only go through the TUN. Windows sends a name
// query to the DNS servers of all interfaces, not only to those of the TUN:
// to the first server of each interface, then to all of them when no answer
// arrives within a second or two. It sends the queries for the servers of
// an interface out through that interface, whatever the routes say, and
// other programs reach an on-link resolver, like 192.168.1.1 from DHCP,
// through its LAN route, which is more specific than the TUN's default
// route. Since Windows 11 and Server 2022, Windows may also send its
// queries over HTTPS or TLS, so there its DNS Client service may not
// connect outside the TUN at all, except for name resolution on the local
// link (mDNS, LLMNR).
// - ipv4, ipv6: no IPv4, or no IPv6, at all, in either direction, for a TUN
// that no route of it leads to, except loopback and what Windows itself
// needs on the local link (DHCP, and for IPv6 neighbor and multicast
// listener discovery), none of which can leave it. The TUN carries what
// is routed to it even without an address of that IP version in gateway:
// Windows gives it link-local ones itself, an IPv6 one at once, an IPv4
// one from 169.254.0.0/16 after some seconds (until then, IPv4 routed to
// the TUN is unreachable).
//
// The filters live in a dynamic WFP session: closing the returned engine handle
// with closeWFPEngine deletes them, and so does Windows when the process dies.
func blockLeaks(tun winipcfg.LUID, dns, ipv4, ipv6 bool) (windows.Handle, error) {
engine, err := openWFPEngine()
if err != nil {
return 0, err
}
if err := fwpmResult(procFwpmTransactionBegin0.Call(uintptr(engine), 0)); err != nil {
closeWFPEngine(engine)
return 0, errors.New("FwpmTransactionBegin0 failed").Base(err)
}
err = addLeakFilters(engine, tun, dns, ipv4, ipv6)
if err == nil {
if err = fwpmResult(procFwpmTransactionCommit0.Call(uintptr(engine))); err != nil {
err = errors.New("FwpmTransactionCommit0 failed").Base(err)
}
}
if err != nil {
procFwpmTransactionAbort0.Call(uintptr(engine))
closeWFPEngine(engine)
return 0, err
}
return engine, nil
}
func openWFPEngine() (windows.Handle, error) {
if err := modfwpuclnt.Load(); err != nil {
return 0, err
}
// txnWaitTimeoutInMSec stays 0 for BFE's default, so that a transaction
// held by another program cannot hang the start forever.
session := fwpmSession0{
displayData: fwpmDisplayData0{name: utf16Ptr("Xray TUN")},
flags: fwpmSessionFlagDynamic,
}
var engine windows.Handle
if err := fwpmResult(procFwpmEngineOpen0.Call(0, rpcCAuthnWinNT, 0, uintptr(unsafe.Pointer(&session)), uintptr(unsafe.Pointer(&engine)))); err != nil {
return 0, errors.New("FwpmEngineOpen0 failed").Base(err)
}
return engine, nil
}
func closeWFPEngine(engine windows.Handle) {
procFwpmEngineClose0.Call(uintptr(engine))
}
// addLeakFilters adds the filters of blockLeaks in a sublayer of their own.
// blockLeaks runs it in a transaction, so that they take effect all at once.
func addLeakFilters(engine windows.Handle, tun winipcfg.LUID, dns, ipv4, ipv6 bool) error {
exe, err := os.Executable()
if err != nil {
return err
}
exePath, err := windows.UTF16PtrFromString(exe)
if err != nil {
return err
}
var appID *fwpByteBlob
if err := fwpmResult(procFwpmGetAppIdFromFileName0.Call(uintptr(unsafe.Pointer(exePath)), uintptr(unsafe.Pointer(&appID)))); err != nil {
return errors.New("FwpmGetAppIdFromFileName0 failed for ", exe).Base(err)
}
defer func() { procFwpmFreeMemory0.Call(uintptr(unsafe.Pointer(&appID))) }()
sublayer := fwpmSublayer0{
displayData: fwpmDisplayData0{name: utf16Ptr("Xray TUN")},
weight: 0xffff,
}
if sublayer.subLayerKey, err = windows.GenerateGUID(); err != nil {
return err
}
if err := fwpmResult(procFwpmSubLayerAdd0.Call(uintptr(engine), uintptr(unsafe.Pointer(&sublayer)), 0)); err != nil {
return errors.New("FwpmSubLayerAdd0 failed").Base(err)
}
add := func(layer *windows.GUID, name string, flags, action uint32, weight uint8, conditions ...fwpmFilterCondition0) error {
return addFilter(engine, &sublayer.subLayerKey, layer, "Xray TUN: "+name, flags, action, weight, conditions...)
}
var pinner runtime.Pinner
defer pinner.Unpin()
tunLUID := new(uint64)
*tunLUID = uint64(tun)
pinner.Pin(tunLUID) // the condition only holds it as uintptr
// The heaviest matching filter of a sublayer decides. All sublayers have
// their say, though, and a block in any of them beats a permit, unless
// the permit is hard: it clears the action right, and then the blocks of
// lower sublayers, Windows Firewall rules among them, no longer override
// it, only a callout's veto does. Xray's own connections out get such a
// hard permit. Connections from outside to Xray get an ordinary one, so
// that firewalls keep guarding its inbounds.
self := condition(&fwpmConditionALEAppID, fwpByteBlobType, uintptr(unsafe.Pointer(appID)))
dns53 := condition(&fwpmConditionIPRemotePort, fwpUint16, 53)
// DNS goes through the TUN when its local address is the TUN's, and it
// also leaves, or arrives, through the TUN. The local address alone
// decides by default, but with weak host sending or receiving enabled,
// packets of the TUN's address can use other interfaces. (The next hop,
// the interface replies would leave by, is not known for arriving ones.)
onTUN := func(field *windows.GUID) fwpmFilterCondition0 {
return condition(field, fwpUint64, uintptr(unsafe.Pointer(tunLUID)))
}
out := []fwpmFilterCondition0{dns53, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPNexthopInterface)}
in := []fwpmFilterCondition0{dns53, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPArrivalInterface)}
for _, layer := range []struct {
key *windows.GUID
selfFlags uint32
throughTUN []fwpmFilterCondition0
}{
{&fwpmLayerALEAuthConnectV4, fwpmFilterFlagClearActionRight, out},
{&fwpmLayerALEAuthRecvAcceptV4, 0, in},
{&fwpmLayerALEAuthConnectV6, fwpmFilterFlagClearActionRight, out},
{&fwpmLayerALEAuthRecvAcceptV6, 0, in},
} {
if err := add(layer.key, "permit Xray", layer.selfFlags, fwpActionPermit, 4, self); err != nil {
return err
}
if dns {
if err := add(layer.key, "permit DNS through the TUN", 0, fwpActionPermit, 3, layer.throughTUN...); err != nil {
return err
}
if err := add(layer.key, "block DNS", 0, fwpActionBlock, 2, dns53); err != nil {
return err
}
}
}
// Since Windows 11 and Server 2022 (build 20348), the DNS Client service
// may also send the queries for an interface's servers over HTTPS or TLS,
// out through that interface and to any port. So there it may only
// connect through the TUN, except for mDNS and LLMNR, which stay on the
// local link (over an IP version only while it is not blocked altogether).
// Earlier versions only query port 53, and may run the service in one
// process with others, which the filters would catch as well. Like
// Windows Firewall's rules for it, they recognize the service by its SID,
// which Windows puts in the token of its process: the security descriptor
// grants that SID the right to match (FWP_ACTRL_MATCH_FILTER, CC in SDDL).
if _, _, build := windows.RtlGetNtVersionNumbers(); dns && build >= 20348 {
sd, err := windows.SecurityDescriptorFromString("O:SYG:SYD:(A;;CCRC;;;" + dnsClientSID + ")")
if err != nil {
return err
}
sdBlob := &fwpByteBlob{size: sd.Length(), data: (*byte)(unsafe.Pointer(sd))}
pinner.Pin(sdBlob) // the condition only holds it as uintptr
dnsClient := condition(&fwpmConditionALEUserID, fwpSecurityDescriptorType, uintptr(unsafe.Pointer(sdBlob)))
// Conditions on the same field match when any of them does.
mdnsLLMNR := []fwpmFilterCondition0{dnsClient, condition(&fwpmConditionIPRemotePort, fwpUint16, 5353), condition(&fwpmConditionIPRemotePort, fwpUint16, 5355)}
for _, layer := range []struct {
key *windows.GUID
localLink bool
}{
{&fwpmLayerALEAuthConnectV4, !ipv4},
{&fwpmLayerALEAuthConnectV6, !ipv6},
} {
if err := add(layer.key, "permit the DNS Client service through the TUN", 0, fwpActionPermit, 3, dnsClient, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPNexthopInterface)); err != nil {
return err
}
if layer.localLink {
if err := add(layer.key, "permit the DNS Client service's mDNS and LLMNR", 0, fwpActionPermit, 3, mdnsLLMNR...); err != nil {
return err
}
}
if err := add(layer.key, "block the DNS Client service", 0, fwpActionBlock, 2, dnsClient); err != nil {
return err
}
}
}
// Both directions: replies to a connection accepted from outside would
// leave through the physical link as well.
loopback := fwpmFilterCondition0{
fieldKey: fwpmConditionFlags,
matchType: fwpMatchFlagsAllSet,
conditionValue: fwpValue0{typ: fwpUint32, value: fwpConditionFlagIsLoopback},
}
if ipv4 {
// DHCP keeps the addresses of the other interfaces, which Xray's own
// connections use.
dhcp := []fwpmFilterCondition0{
condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_UDP),
condition(&fwpmConditionIPLocalPort, fwpUint16, 68),
condition(&fwpmConditionIPRemotePort, fwpUint16, 67),
}
for _, layer := range []*windows.GUID{&fwpmLayerALEAuthConnectV4, &fwpmLayerALEAuthRecvAcceptV4} {
if err := add(layer, "permit IPv4 loopback", 0, fwpActionPermit, 1, loopback); err != nil {
return err
}
if err := add(layer, "permit DHCP", 0, fwpActionPermit, 1, dhcp...); err != nil {
return err
}
if err := add(layer, "block IPv4", 0, fwpActionBlock, 0); err != nil {
return err
}
}
}
if ipv6 {
// Neighbor and multicast listener discovery, ICMPv6 130-137 and 143,
// whose type and code sit where the local and remote port are.
discovery := []fwpmFilterCondition0{condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_ICMPV6)}
for _, typ := range []uintptr{130, 131, 132, 133, 134, 135, 136, 137, 143} {
discovery = append(discovery, condition(&fwpmConditionIPLocalPort, fwpUint16, typ))
}
discovery = append(discovery, condition(&fwpmConditionIPRemotePort, fwpUint16, 0))
dhcpv6 := []fwpmFilterCondition0{
condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_UDP),
condition(&fwpmConditionIPLocalPort, fwpUint16, 546),
condition(&fwpmConditionIPRemotePort, fwpUint16, 547),
}
for _, direction := range []struct {
layer *windows.GUID
dhcpv6 []fwpmFilterCondition0
}{
// The client sends to the servers' multicast address, and they
// answer from their own.
{&fwpmLayerALEAuthConnectV6, slices.Concat(dhcpv6, []fwpmFilterCondition0{condition(&fwpmConditionIPRemoteAddress, fwpByteArray16Type, uintptr(unsafe.Pointer(&ipv6AllDHCPv6Servers)))})},
{&fwpmLayerALEAuthRecvAcceptV6, dhcpv6},
} {
if err := add(direction.layer, "permit IPv6 loopback", 0, fwpActionPermit, 1, loopback); err != nil {
return err
}
if err := add(direction.layer, "permit IPv6 neighbor and multicast listener discovery", 0, fwpActionPermit, 1, discovery...); err != nil {
return err
}
if err := add(direction.layer, "permit DHCPv6", 0, fwpActionPermit, 1, direction.dhcpv6...); err != nil {
return err
}
if err := add(direction.layer, "block IPv6", 0, fwpActionBlock, 0); err != nil {
return err
}
}
}
return nil
}
func addFilter(engine windows.Handle, sublayer, layer *windows.GUID, name string, flags, action uint32, weight uint8, conditions ...fwpmFilterCondition0) error {
filter := fwpmFilter0{
displayData: fwpmDisplayData0{name: utf16Ptr(name)},
flags: flags,
layerKey: *layer,
subLayerKey: *sublayer,
weight: fwpValue0{typ: fwpUint8, value: uintptr(weight)},
numFilterConditions: uint32(len(conditions)),
action: fwpmAction0{typ: action},
}
if len(conditions) > 0 {
filter.filterCondition = &conditions[0]
}
if err := fwpmResult(procFwpmFilterAdd0.Call(uintptr(engine), uintptr(unsafe.Pointer(&filter)), 0, 0)); err != nil {
return errors.New("FwpmFilterAdd0 failed for ", name).Base(err)
}
return nil
}
// dnsOutsideTUN returns the servers outside all of prefixes, the TUN's own
// subnets and routes: queries to them cannot go through the TUN.
func dnsOutsideTUN(servers []netip.Addr, prefixes []netip.Prefix) []netip.Addr {
var outside []netip.Addr
for _, server := range servers {
server = server.Unmap()
if !slices.ContainsFunc(prefixes, func(p netip.Prefix) bool { return p.Contains(server) }) {
outside = append(outside, server)
}
}
return outside
}
// flushDNSCache drops the answers Windows cached so far, like ipconfig
// /flushdns, so that names get resolved again with the current DNS setup.
func flushDNSCache() error {
if err := procDnsFlushResolverCache.Find(); err != nil {
return err
}
if r, _, err := procDnsFlushResolverCache.Call(); r == 0 {
return err
}
return nil
}
+206
View File
@@ -0,0 +1,206 @@
//go:build windows
package tun
import (
"context"
go_errors "errors"
"net"
"net/netip"
"slices"
"testing"
"unsafe"
"github.com/xtls/xray-core/transport/internet"
"golang.org/x/sys/windows"
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
)
// The WFP structures are handed to fwpuclnt.dll as they are, so their layout
// has to match what MSVC produces for 64-bit and for 32-bit Windows.
func TestWFPStructLayout(t *testing.T) {
check := func(name string, got, want64, want32 []uintptr) {
t.Helper()
want := want32
if unsafe.Sizeof(uintptr(0)) == 8 {
want = want64
}
if !slices.Equal(got, want) {
t.Errorf("%s: size and offsets are %v, want %v", name, got, want)
}
}
var blob fwpByteBlob
check("FWP_BYTE_BLOB",
[]uintptr{unsafe.Sizeof(blob), unsafe.Offsetof(blob.data)},
[]uintptr{16, 8}, []uintptr{8, 4})
var value fwpValue0
check("FWP_VALUE0",
[]uintptr{unsafe.Sizeof(value), unsafe.Offsetof(value.value)},
[]uintptr{16, 8}, []uintptr{8, 4})
var display fwpmDisplayData0
check("FWPM_DISPLAY_DATA0",
[]uintptr{unsafe.Sizeof(display), unsafe.Offsetof(display.description)},
[]uintptr{16, 8}, []uintptr{8, 4})
var action fwpmAction0
check("FWPM_ACTION0",
[]uintptr{unsafe.Sizeof(action), unsafe.Offsetof(action.filterType)},
[]uintptr{20, 4}, []uintptr{20, 4})
var cond fwpmFilterCondition0
check("FWPM_FILTER_CONDITION0",
[]uintptr{unsafe.Sizeof(cond), unsafe.Offsetof(cond.matchType), unsafe.Offsetof(cond.conditionValue)},
[]uintptr{40, 16, 24}, []uintptr{28, 16, 20})
var session fwpmSession0
check("FWPM_SESSION0",
[]uintptr{
unsafe.Sizeof(session), unsafe.Offsetof(session.displayData), unsafe.Offsetof(session.flags),
unsafe.Offsetof(session.txnWaitTimeoutInMSec), unsafe.Offsetof(session.processID), unsafe.Offsetof(session.sid),
unsafe.Offsetof(session.username), unsafe.Offsetof(session.kernelMode),
},
[]uintptr{72, 16, 32, 36, 40, 48, 56, 64},
[]uintptr{48, 16, 24, 28, 32, 36, 40, 44})
var sublayer fwpmSublayer0
check("FWPM_SUBLAYER0",
[]uintptr{
unsafe.Sizeof(sublayer), unsafe.Offsetof(sublayer.displayData), unsafe.Offsetof(sublayer.flags),
unsafe.Offsetof(sublayer.providerKey), unsafe.Offsetof(sublayer.providerData), unsafe.Offsetof(sublayer.weight),
},
[]uintptr{72, 16, 32, 40, 48, 64},
[]uintptr{44, 16, 24, 28, 32, 40})
var filter fwpmFilter0
check("FWPM_FILTER0",
[]uintptr{
unsafe.Sizeof(filter), unsafe.Offsetof(filter.displayData), unsafe.Offsetof(filter.flags),
unsafe.Offsetof(filter.providerKey), unsafe.Offsetof(filter.providerData), unsafe.Offsetof(filter.layerKey),
unsafe.Offsetof(filter.subLayerKey), unsafe.Offsetof(filter.weight), unsafe.Offsetof(filter.numFilterConditions),
unsafe.Offsetof(filter.filterCondition), unsafe.Offsetof(filter.action), unsafe.Offsetof(filter.providerContextKey),
unsafe.Offsetof(filter.reserved), unsafe.Offsetof(filter.filterID), unsafe.Offsetof(filter.effectiveWeight),
},
[]uintptr{200, 16, 32, 40, 48, 64, 80, 96, 112, 120, 128, 152, 168, 176, 184},
[]uintptr{152, 16, 24, 28, 32, 40, 56, 72, 80, 84, 88, 112, 128, 136, 144})
}
// TestLeakFiltersAccepted has WFP validate the filters by adding them inside a
// transaction that is then aborted, which leaves the system untouched. Adding
// filters requires an elevated process.
func TestLeakFiltersAccepted(t *testing.T) {
skipUnlessElevated := func(err error) {
t.Helper()
if go_errors.Is(err, windows.ERROR_ACCESS_DENIED) {
t.Skipf("WFP filters can only be added by an elevated process: %v", err)
}
t.Fatal(err)
}
engine, err := openWFPEngine()
if err != nil {
skipUnlessElevated(err)
}
defer closeWFPEngine(engine)
if err := fwpmResult(procFwpmTransactionBegin0.Call(uintptr(engine), 0)); err != nil {
skipUnlessElevated(err)
}
defer procFwpmTransactionAbort0.Call(uintptr(engine))
// Any interface stands in for the TUN; the loopback one always exists.
loopback, err := winipcfg.LUIDFromIndex(1)
if err != nil {
t.Fatal(err)
}
if err := addLeakFilters(engine, loopback, true, true, true); err != nil {
skipUnlessElevated(err)
}
}
func TestDNSClientSID(t *testing.T) {
sid, _, _, err := windows.LookupSID("", `NT SERVICE\Dnscache`)
if err != nil {
t.Fatal(err)
}
if sid.String() != dnsClientSID {
t.Errorf(`NT SERVICE\Dnscache is %v, not %v`, sid, dnsClientSID)
}
}
func TestDNSOutsideTUN(t *testing.T) {
prefixes := []netip.Prefix{
netip.MustParsePrefix("198.51.100.1/30"), // gateway, not masked
netip.MustParsePrefix("203.0.113.0/24"), // route
}
servers := []netip.Addr{
netip.MustParseAddr("198.51.100.2"),
netip.MustParseAddr("203.0.113.53"),
netip.MustParseAddr("::ffff:203.0.113.54"),
netip.MustParseAddr("8.8.8.8"),
netip.MustParseAddr("2001:db8::53"),
}
want := []netip.Addr{netip.MustParseAddr("8.8.8.8"), netip.MustParseAddr("2001:db8::53")}
if got := dnsOutsideTUN(servers, prefixes); !slices.Equal(got, want) {
t.Errorf("got %v, want %v", got, want)
}
}
func TestResolveOnOwn(t *testing.T) {
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("::ffff:203.0.113.53")})
t.Cleanup(func() { internet.SkipDNSServers(nil) })
preferGo, dial := net.DefaultResolver.PreferGo, net.DefaultResolver.Dial
saved := resolveOnOwn()
t.Cleanup(saved.restore)
if !net.DefaultResolver.PreferGo || net.DefaultResolver.Dial == nil {
t.Fatal("net.DefaultResolver is unchanged")
}
if _, err := net.DefaultResolver.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil {
t.Error("the TUN's DNS server was not skipped")
}
conn, err := net.DefaultResolver.Dial(context.Background(), "udp", "127.0.0.1:53")
if err != nil {
t.Fatal(err)
}
conn.Close()
saved.restore()
if net.DefaultResolver.PreferGo != preferGo || (net.DefaultResolver.Dial == nil) != (dial == nil) {
t.Error("net.DefaultResolver is not restored")
}
}
// TestTunOnlyDNS checks that a DNS server another interface uses as well is
// not skipped, while one of the TUN alone is.
func TestTunOnlyDNS(t *testing.T) {
adapters, err := winipcfg.GetAdaptersAddresses(windows.AF_UNSPEC, winipcfg.GAAFlagIncludeGateways)
if err != nil {
t.Fatal(err)
}
var other netip.Addr
for _, adapter := range adapters {
if adapter.OperStatus == winipcfg.IfOperStatusUp && adapter.FirstGatewayAddress != nil && adapter.FirstDNSServerAddress != nil {
other, _ = netip.AddrFromSlice(adapter.FirstDNSServerAddress.Address.IP())
other = other.Unmap()
break
}
}
if !other.IsValid() {
t.Skip("no interface with a gateway and a DNS server")
}
tunOnly := netip.MustParseAddr("203.0.113.53")
// LUID 0 is no interface, so every one counts as another.
got, err := tunOnlyDNS(0, []netip.Addr{other, tunOnly})
if err != nil {
t.Fatal(err)
}
if !slices.Equal(got, []netip.Addr{tunOnly}) {
t.Errorf("got %v, want [%v]", got, tunOnly)
}
}
func TestFlushDNSCache(t *testing.T) {
if err := flushDNSCache(); err != nil {
t.Fatal(err)
}
}
+32
View File
@@ -0,0 +1,32 @@
package internet
import (
"net/netip"
"slices"
"sync/atomic"
)
var skippedDNSServers atomic.Pointer[[]netip.Addr]
// SkipDNSServers has the queries Xray sends to the system's DNS servers on its
// own, like those of localdns, skip servers until it is called again. The DNS
// servers of a TUN are only meant for what goes through it: queried by Xray
// itself they lead back into it, or nowhere.
func SkipDNSServers(servers []netip.Addr) {
skipped := make([]netip.Addr, len(servers))
for i, server := range servers {
skipped[i] = server.Unmap()
}
skippedDNSServers.Store(&skipped)
}
// IsSkippedDNSServer reports whether address, a DNS server as host:port, is to
// be skipped, see SkipDNSServers.
func IsSkippedDNSServer(address string) bool {
skipped := skippedDNSServers.Load()
if skipped == nil {
return false
}
server, err := netip.ParseAddrPort(address)
return err == nil && slices.Contains(*skipped, server.Addr().Unmap())
}
+27
View File
@@ -0,0 +1,27 @@
package internet_test
import (
"net/netip"
"testing"
"github.com/xtls/xray-core/transport/internet"
)
func TestSkipDNSServers(t *testing.T) {
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("::ffff:203.0.113.53"), netip.MustParseAddr("2001:db8::53")})
t.Cleanup(func() { internet.SkipDNSServers(nil) })
for address, want := range map[string]bool{
"203.0.113.53:53": true,
"[2001:db8::53]:53": true,
"198.51.100.53:53": false,
"localhost:53": false,
} {
if got := internet.IsSkippedDNSServer(address); got != want {
t.Errorf("IsSkippedDNSServer(%q) = %v, want %v", address, got, want)
}
}
internet.SkipDNSServers(nil)
if internet.IsSkippedDNSServer("203.0.113.53:53") {
t.Error("still skipped after SkipDNSServers(nil)")
}
}
+4 -29
View File
@@ -18,7 +18,6 @@ import (
"regexp"
"strings"
"sync"
"sync/atomic"
"time"
"unsafe"
@@ -37,18 +36,6 @@ import (
type Conn struct {
*reality.Conn
suppressCloseNotify atomic.Bool
}
func (c *Conn) SuppressCloseNotify() {
c.suppressCloseNotify.Store(true)
}
func (c *Conn) Close() error {
if c.suppressCloseNotify.Load() {
return c.Conn.NetConn().Close()
}
return c.Conn.Close()
}
func (c *Conn) HandshakeAddress() net.Address {
@@ -69,22 +56,10 @@ func Server(c net.Conn, config *reality.Config) (net.Conn, error) {
type UConn struct {
*utls.UConn
Config *Config
ServerName string
AuthKey []byte
Verified bool
suppressCloseNotify atomic.Bool
}
func (c *UConn) SuppressCloseNotify() {
c.suppressCloseNotify.Store(true)
}
func (c *UConn) Close() error {
if c.suppressCloseNotify.Load() {
return c.NetConn().Close()
}
return c.UConn.Close()
Config *Config
ServerName string
AuthKey []byte
Verified bool
}
func (c *UConn) HandshakeAddress() net.Address {
-17
View File
@@ -6,7 +6,6 @@ import (
"crypto/tls"
"math/big"
"slices"
"sync/atomic"
"time"
utls "github.com/refraction-networking/utls"
@@ -30,19 +29,11 @@ var (
type Conn struct {
*tls.Conn
suppressCloseNotify atomic.Bool
}
const tlsCloseTimeout = 250 * time.Millisecond
func (c *Conn) SuppressCloseNotify() {
c.suppressCloseNotify.Store(true)
}
func (c *Conn) Close() error {
if c.suppressCloseNotify.Load() {
return c.Conn.NetConn().Close()
}
timer := time.AfterFunc(tlsCloseTimeout, func() {
c.Conn.NetConn().Close()
})
@@ -83,19 +74,11 @@ func Server(c net.Conn, config *tls.Config) net.Conn {
type UConn struct {
*utls.UConn
suppressCloseNotify atomic.Bool
}
var _ Interface = (*UConn)(nil)
func (c *UConn) SuppressCloseNotify() {
c.suppressCloseNotify.Store(true)
}
func (c *UConn) Close() error {
if c.suppressCloseNotify.Load() {
return c.Conn.NetConn().Close()
}
timer := time.AfterFunc(tlsCloseTimeout, func() {
c.Conn.NetConn().Close()
})