Compare commits

...
Author SHA1 Message Date
Esko Mobius 7d3e44fee2 Proxy: Add MASQUE outbound & transport (IETF CONNECT-IP, RFC 9484) (#6807)
Closes https://github.com/XTLS/Xray-core/issues/5495#issuecomment-3710683679
2026-09-24 02:13:37 +00:00
dependabot[bot] 9927942aaa Bump google.golang.org/grpc from 1.83.2 to 1.84.0 (#6793)
Bumps [google.golang.org/grpc](https://github.com/grpc/grpc-go) from 1.83.2 to 1.84.0.
- [Release notes](https://github.com/grpc/grpc-go/releases)
- [Commits](https://github.com/grpc/grpc-go/compare/v1.83.2...v1.84.0)

---
updated-dependencies:
- dependency-name: google.golang.org/grpc
  dependency-version: 1.84.0
  dependency-type: direct:production
  update-type: version-update:semver-minor
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-09-24 02:04:39 +00:00
LjhAUMEM a308ded2e6 WireGuard outbound: Fix endpoint IP address (#6804)
Fixes https://github.com/XTLS/Xray-core/issues/6803
2026-09-24 01:59:28 +00:00
dependabot[bot] 7741e9e77e Bump golang.zx2c4.com/wireguard/windows from 1.0.1 to 1.1.1 (#6809)
Bumps golang.zx2c4.com/wireguard/windows from 1.0.1 to 1.1.1.

---
updated-dependencies:
- dependency-name: golang.zx2c4.com/wireguard/windows
  dependency-version: 1.1.1
  dependency-type: direct:production
  update-type: version-update:semver-minor
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-09-24 01:57:16 +00:00
39 changed files with 5198 additions and 37 deletions
+3 -3
View File
@@ -31,8 +31,8 @@ require (
golang.org/x/sys v0.48.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
golang.zx2c4.com/wireguard/windows v1.0.1
google.golang.org/grpc v1.83.2
golang.zx2c4.com/wireguard/windows v1.1.1
google.golang.org/grpc v1.84.0
google.golang.org/protobuf v1.36.12
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
h12.io/socks v1.0.3
@@ -60,6 +60,6 @@ require (
golang.org/x/text v0.42.0 // indirect
golang.org/x/time v0.14.0 // indirect
golang.org/x/tools v0.49.0 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 // indirect
gopkg.in/yaml.v2 v2.4.0 // indirect
)
+6 -24
View File
@@ -2,16 +2,10 @@ github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sx
github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig=
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344/go.mod h1:GIjDIg/heH5DOkXY3YJ/wNhfHsQHoXGjl8G8amsYQ1I=
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
github.com/go-quicktest/qt v1.102.0 h1:HSQxCeh5YZH3EL3W39ixjtyaEhcWSXQHtHnMBzSs474=
github.com/go-quicktest/qt v1.102.0/go.mod h1:p4lGIVX+8Wa6ZPNDvqcxq36XpUDLh42FLetFU7odllI=
github.com/golang/mock v1.7.0-rc.1 h1:YojYx61/OLFsiv6Rw1Z96LpldJIy31o+UHmwAUMJ6/U=
@@ -91,18 +85,6 @@ github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguH
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU=
go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc=
go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc=
go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo=
go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58=
go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0=
go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI=
go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA=
go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk=
go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE=
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
@@ -157,14 +139,14 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8=
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
golang.zx2c4.com/wireguard/windows v1.1.1 h1:8/H97U1v1PNDNcBsMZgU3KFuND9MQdTsU2NOwmCXArE=
golang.zx2c4.com/wireguard/windows v1.1.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU=
google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 h1:qEHAMpSaUhtD0p3NbEEI83HwNGFxEwaSJ1G9PLnCBZE=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.84.0 h1:soMyaPJ8pAak5PIQ0DGBUir0XRo2fRoMqhNWMLlLxO0=
google.golang.org/grpc v1.84.0/go.mod h1:ljCht0DrxQrXBDRTZp52Qxh3Ffk8CdYm2sj4O2QN2C0=
google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc=
google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
+37
View File
@@ -0,0 +1,37 @@
package conf
import (
"net/netip"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/proxy/masque"
"google.golang.org/protobuf/proto"
)
type MasqueClientConfig struct {
Address *Address `json:"address"`
Port uint16 `json:"port"`
RemoteDNS []string `json:"remoteDNS"`
}
func (c *MasqueClientConfig) Build() (proto.Message, error) {
if c.Address == nil {
return nil, errors.New(`MASQUE: "address" is not set`)
}
if c.Port == 0 {
return nil, errors.New(`MASQUE: "port" is not set`)
}
for _, s := range c.RemoteDNS {
if _, err := netip.ParseAddr(s); err != nil {
return nil, errors.New(`MASQUE: invalid "remoteDNS" `, s).Base(err)
}
}
return &masque.ClientConfig{
Server: &protocol.ServerEndpoint{
Address: c.Address.Build(),
Port: uint32(c.Port),
},
RemoteDns: c.RemoteDNS,
}, nil
}
+85
View File
@@ -0,0 +1,85 @@
package conf_test
import (
"encoding/json"
"testing"
. "github.com/xtls/xray-core/infra/conf"
"github.com/xtls/xray-core/transport/internet/masque"
)
func TestMasqueConfig(t *testing.T) {
creator := func() Buildable {
return new(MasqueConfig)
}
runMultiTestCase(t, []TestCase{
{
Input: `{}`,
Parser: loadJSON(creator),
Output: &masque.Config{Path: "/.well-known/masque/ip/*/*/"},
},
{
Input: `{
"host": "example.com:8443",
"path": "/.well-known/masque/ip/{target}/{ipproto}/",
"headers": {"Authorization": "Basic dTpw"}
}`,
Parser: loadJSON(creator),
Output: &masque.Config{
Host: "example.com:8443",
Path: "/.well-known/masque/ip/*/*/",
Headers: map[string]string{"Authorization": "Basic dTpw"},
},
},
{
Input: `{"path": "/masque/ip{?target,ipproto}"}`,
Parser: loadJSON(creator),
Output: &masque.Config{Path: "/masque/ip?target=*&ipproto=*"},
},
})
for _, input := range []string{
`{"path": "/masque/{target}/{ipproto}/{dns}"}`,
`{"path": "masque"}`,
`{"host": "example.com/path"}`,
`{"headers": {"host": "example.com"}}`,
`{"headers": {"Capsule-Protocol": "?0"}}`,
`{"headers": {"X Token": "a"}}`,
`{"headers": {"X-Token": "a\r\nb"}}`,
} {
if _, err := loadJSON(creator)(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
func TestMasqueOutboundConfig(t *testing.T) {
build := func(s string) error {
c := new(OutboundDetourConfig)
if err := json.Unmarshal([]byte(s), c); err != nil {
return err
}
_, err := c.Build()
return err
}
if err := build(`{
"protocol": "masque",
"settings": {"address": "example.com", "port": 443},
"streamSettings": {"network": "masque", "security": "tls"},
"mux": {"enabled": false, "concurrency": -1}
}`); err != nil {
t.Error(err)
}
for _, input := range []string{
`{"protocol": "masque", "settings": {"address": "example.com"}, "streamSettings": {"network": "masque", "security": "tls"}}`,
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true}}`,
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true, "concurrency": -1}}`,
`{"protocol": "freedom", "streamSettings": {"network": "masque", "security": "tls"}}`,
} {
if err := build(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
+13
View File
@@ -36,6 +36,8 @@ func (p TransportProtocol) Build() (string, error) {
return "", errors.PrintRemovedFeatureError("QUIC transport (without web service, etc.)", "XHTTP stream-one H3")
case "hysteria":
return "hysteria", nil
case "masque":
return "masque", nil
case "xdrive":
return "xdrive", nil
default:
@@ -61,6 +63,7 @@ type StreamConfig struct {
WSSettings *WebSocketConfig `json:"wsSettings"`
HTTPUPGRADESettings *HttpUpgradeConfig `json:"httpupgradeSettings"`
HysteriaSettings *HysteriaConfig `json:"hysteriaSettings"`
MASQUESettings *MasqueConfig `json:"masqueSettings"`
XDRIVESettings *XDriveConfig `json:"xdriveSettings"`
SocketSettings *SocketConfig `json:"sockopt"`
}
@@ -195,6 +198,16 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
Settings: serial.ToTypedMessage(hs),
})
}
if c.MASQUESettings != nil {
ms, err := c.MASQUESettings.Build()
if err != nil {
return nil, errors.New("Failed to build MASQUE config.").Base(err)
}
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
ProtocolName: "masque",
Settings: serial.ToTypedMessage(ms),
})
}
if c.XDRIVESettings != nil {
xs, err := c.XDRIVESettings.Build()
if err != nil {
+42
View File
@@ -20,10 +20,12 @@ import (
"github.com/xtls/xray-core/transport/internet/httpupgrade"
"github.com/xtls/xray-core/transport/internet/hysteria"
"github.com/xtls/xray-core/transport/internet/kcp"
"github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/splithttp"
"github.com/xtls/xray-core/transport/internet/tcp"
"github.com/xtls/xray-core/transport/internet/websocket"
"github.com/xtls/xray-core/transport/internet/xdrive"
"golang.org/x/net/http/httpguts"
"google.golang.org/protobuf/proto"
)
@@ -786,6 +788,46 @@ func (c *HysteriaConfig) Build() (proto.Message, error) {
return config, nil
}
type MasqueConfig struct {
Host string `json:"host"`
Path string `json:"path"`
Headers map[string]string `json:"headers"`
}
func (c *MasqueConfig) Build() (proto.Message, error) {
path := c.Path
if path == "" {
path = masque.DefaultPath
}
path = strings.NewReplacer(
"{target}", "*", "{ipproto}", "*",
"{?target,ipproto}", "?target=*&ipproto=*", "{?ipproto,target}", "?ipproto=*&target=*",
"{&target,ipproto}", "&target=*&ipproto=*", "{&ipproto,target}", "&ipproto=*&target=*",
).Replace(path)
if !strings.HasPrefix(path, "/") || strings.ContainsAny(path, "{}") {
return nil, errors.New(`invalid "path": `, path, `, only the variables {target} and {ipproto} are supported`)
}
if c.Host != "" {
if u, err := url.Parse("https://" + c.Host); err != nil || u.Host != c.Host {
return nil, errors.New(`invalid "host": `, c.Host)
}
}
for k, v := range c.Headers {
if !httpguts.ValidHeaderFieldName(k) || !httpguts.ValidHeaderFieldValue(v) {
return nil, errors.New(`invalid header in "headers": `, strconv.Quote(k))
}
switch strings.ToLower(k) {
case "host", "capsule-protocol":
return nil, errors.New(`"headers" can't contain "`, k, `"`)
}
}
return &masque.Config{
Host: c.Host,
Path: path,
Headers: c.Headers,
}, nil
}
func readFileOrString(f string, s []string) ([]byte, error) {
if len(f) > 0 {
return filesystem.ReadCert(f)
+10
View File
@@ -16,6 +16,7 @@ import (
"github.com/xtls/xray-core/common/serial"
core "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/proxy/freedom"
"github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/transport/internet"
)
@@ -48,6 +49,7 @@ var (
"vmess": func() interface{} { return new(VMessOutboundConfig) },
"trojan": func() interface{} { return new(TrojanClientConfig) },
"hysteria": func() interface{} { return new(HysteriaClientConfig) },
"masque": func() interface{} { return new(MasqueClientConfig) },
"dns": func() interface{} { return new(DNSOutboundConfig) },
"wireguard": func() interface{} { return &WireGuardConfig{IsClient: true} },
}, "protocol", "settings")
@@ -338,6 +340,14 @@ func (c *OutboundDetourConfig) Build() (*core.OutboundHandlerConfig, error) {
return nil, err
}
if _, ok := ts.(*masque.ClientConfig); ok {
if ms := senderSettings.MultiplexSettings; ms != nil && ms.Enabled {
return nil, errors.New(`masque outbound does not support "mux"`)
}
} else if senderSettings.StreamSettings != nil && senderSettings.StreamSettings.ProtocolName == "masque" {
return nil, errors.New("the masque transport can only be used by the masque outbound")
}
if fc, ok := ts.(*freedom.Config); ok {
if senderSettings.StreamSettings != nil &&
senderSettings.StreamSettings.SocketSettings != nil &&
+2
View File
@@ -41,6 +41,7 @@ import (
_ "github.com/xtls/xray-core/proxy/freedom"
_ "github.com/xtls/xray-core/proxy/http"
_ "github.com/xtls/xray-core/proxy/loopback"
_ "github.com/xtls/xray-core/proxy/masque"
_ "github.com/xtls/xray-core/proxy/shadowsocks"
_ "github.com/xtls/xray-core/proxy/socks"
_ "github.com/xtls/xray-core/proxy/trojan"
@@ -54,6 +55,7 @@ import (
_ "github.com/xtls/xray-core/transport/internet/grpc"
_ "github.com/xtls/xray-core/transport/internet/httpupgrade"
_ "github.com/xtls/xray-core/transport/internet/kcp"
_ "github.com/xtls/xray-core/transport/internet/masque"
_ "github.com/xtls/xray-core/transport/internet/reality"
_ "github.com/xtls/xray-core/transport/internet/splithttp"
_ "github.com/xtls/xray-core/transport/internet/tcp"
+328
View File
@@ -0,0 +1,328 @@
package masque
import (
"context"
go_errors "errors"
"io"
"net/netip"
"slices"
"sync"
"sync/atomic"
"time"
"golang.zx2c4.com/wireguard/tun"
"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/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/core"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/proxy/wireguard"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/stat"
"github.com/xtls/xray-core/transport/internet/tls"
)
const (
establishTimeout = 10 * time.Second
retryInterval = time.Second
)
type Client struct {
server *protocol.ServerSpec
policyManager policy.Manager
remoteDNS []netip.Addr
ctx context.Context
cancel context.CancelFunc
tunnel atomic.Pointer[tunnel]
mu sync.Mutex
lastErr error
lastErrAt time.Time
}
func NewClient(ctx context.Context, config *ClientConfig) (*Client, error) {
v := core.MustFromContext(ctx)
p := v.GetFeature(policy.ManagerType()).(policy.Manager)
streamSettings := session.StreamSettingsFromContext(ctx).(*internet.MemoryStreamConfig)
if _, ok := streamSettings.ProtocolSettings.(*masque.Config); !ok {
return nil, errors.New("not masque transport")
}
if tls.ConfigFromStreamSettings(streamSettings) == nil {
return nil, errors.New(`MASQUE requires "security": "tls"`)
}
if config.Server == nil {
return nil, errors.New(`no target server found`)
}
server, err := protocol.NewServerSpecFromPB(config.Server)
if err != nil {
return nil, errors.New("failed to get server spec").Base(err)
}
dns := config.RemoteDns
if len(dns) == 0 {
dns = []string{"1.1.1.1", "1.0.0.1", "2606:4700:4700::1111", "2606:4700:4700::1001"}
}
remoteDNS := make([]netip.Addr, 0, len(dns))
for _, s := range dns {
addr, err := netip.ParseAddr(s)
if err != nil {
return nil, errors.New("invalid remote DNS server ", s).Base(err)
}
remoteDNS = append(remoteDNS, addr)
}
c := &Client{
server: server,
policyManager: p,
remoteDNS: remoteDNS,
}
c.ctx, c.cancel = context.WithCancel(context.Background())
return c, nil
}
func (c *Client) Process(ctx context.Context, link *transport.Link, dialer internet.Dialer) error {
outbounds := session.OutboundsFromContext(ctx)
ob := outbounds[len(outbounds)-1]
if !ob.Target.IsValid() {
return errors.New("target not specified")
}
ob.Name = "masque"
ob.CanSpliceCopy = 3
t, err := c.getTunnel(ctx, dialer)
if err != nil {
return errors.New("failed to establish CONNECT-IP tunnel").Base(err)
}
var newCtx context.Context
var newCancel context.CancelFunc
if session.TimeoutOnlyFromContext(ctx) {
newCtx, newCancel = context.WithCancel(context.Background())
}
sessionPolicy := c.policyManager.ForLevel(0)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, func() {
cancel()
if newCancel != nil {
newCancel()
}
}, sessionPolicy.Timeouts.ConnectionIdle)
if newCtx != nil {
ctx = newCtx
}
var reader buf.Reader
var writer buf.Writer
switch ob.Target.Network {
case net.Network_TCP:
var conn net.Conn
var err error
if sessionPolicy.Timeouts.Handshake != 0 {
timeoutCtx, timeoutCancel := context.WithTimeout(ctx, sessionPolicy.Timeouts.Handshake)
conn, err = t.tnet.DialContext(timeoutCtx, "tcp", ob.Target.NetAddr())
timeoutCancel()
} else {
conn, err = t.tnet.Dial("tcp", ob.Target.NetAddr())
}
if err != nil {
return errors.New("failed to create TCP connection").Base(err)
}
defer conn.Close()
reader = buf.NewReader(conn)
writer = buf.NewWriter(conn)
case net.Network_UDP:
conn, err := t.tnet.Dial("udp", ob.Target.NetAddr())
if err != nil {
return errors.New("failed to create UDP connection").Base(err)
}
defer conn.Close()
uc := &wireguard.UDPConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = uc
writer = uc
default:
panic(ob.Target.Network)
}
requestFunc := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseFunc := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDonePost := task.OnSuccess(responseFunc, task.Close(link.Writer))
if err := task.Run(ctx, requestFunc, responseDonePost); err != nil {
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
return errors.New("connection ends").Base(err)
}
return nil
}
func (c *Client) getTunnel(ctx context.Context, dialer internet.Dialer) (*tunnel, error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.ctx.Err() != nil {
return nil, errors.New("closed")
}
if t := c.tunnel.Load(); t != nil {
select {
case <-t.done:
default:
return t, nil
}
}
if err := ctx.Err(); err != nil {
return nil, err
}
if c.lastErr != nil && time.Since(c.lastErrAt) < retryInterval {
return nil, c.lastErr
}
t, err := c.establish(ctx, dialer)
if err != nil {
c.lastErr, c.lastErrAt = err, time.Now()
return nil, err
}
c.lastErr = nil
c.tunnel.Store(t)
if c.ctx.Err() != nil {
if c.tunnel.CompareAndSwap(t, nil) {
t.close()
}
return nil, errors.New("closed")
}
return t, nil
}
func (c *Client) establish(ctx context.Context, dialer internet.Dialer) (*tunnel, error) {
ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), establishTimeout)
defer cancel()
defer context.AfterFunc(c.ctx, cancel)()
conn, err := dialer.Dial(ctx, c.server.Destination)
if err != nil {
return nil, err
}
mconn, ok := stat.TryUnwrapStatsConn(conn).(*masque.Conn)
if !ok {
conn.Close()
return nil, errors.New("not a CONNECT-IP connection")
}
t, err := newTunnel(conn, mconn.LocalAddrs(), c.remoteDNS)
if err != nil {
conn.Close()
return nil, err
}
errors.LogInfo(ctx, "MASQUE: tunnel established from ", mconn.LocalAddrs())
return t, nil
}
func (c *Client) Close() error {
c.cancel()
if t := c.tunnel.Swap(nil); t != nil {
t.close()
}
return nil
}
type tunnel struct {
conn stat.Connection
dev tun.Device
tnet *wireguard.Net
done chan struct{}
closeOnce sync.Once
}
func newTunnel(conn stat.Connection, local []netip.Addr, remoteDNS []netip.Addr) (*tunnel, error) {
var dns []netip.Addr
for _, addr := range remoteDNS {
if slices.ContainsFunc(local, func(l netip.Addr) bool { return l.Is4() == addr.Is4() }) {
dns = append(dns, addr)
}
}
if len(dns) == 0 {
errors.LogWarning(context.Background(), "MASQUE: no remote DNS server is reachable from the assigned addresses ", local, ", domain names will fail to resolve")
dns = remoteDNS
}
dev, tnet, _, err := wireguard.CreateNetTUN(local, dns, masque.MinPacketSize, true)
if err != nil {
return nil, err
}
t := &tunnel{
conn: conn,
dev: dev,
tnet: tnet,
done: make(chan struct{}),
}
go t.readFromTunnel()
go t.writeToTunnel()
return t, nil
}
func (t *tunnel) readFromTunnel() {
defer t.close()
b := make([]byte, buf.Size)
for {
n, err := t.conn.Read(b)
if err != nil {
if go_errors.Is(err, io.ErrShortBuffer) {
continue
}
errors.LogInfoInner(context.Background(), err, "MASQUE: tunnel closed")
return
}
t.dev.Write([][]byte{b[:n]}, 0)
}
}
func (t *tunnel) writeToTunnel() {
bufs := [][]byte{make([]byte, masque.MinPacketSize)}
sizes := []int{0}
for {
if _, err := t.dev.Read(bufs, sizes, 0); err != nil {
return
}
if _, err := t.conn.Write(bufs[0][:sizes[0]]); err != nil {
var ptb *masque.PacketTooBigError
if go_errors.As(err, &ptb) {
go t.dev.Write([][]byte{ptb.ICMP}, 0)
}
}
}
}
func (t *tunnel) close() {
t.closeOnce.Do(func() {
close(t.done)
t.conn.Close()
t.dev.Close()
})
}
func init() {
common.Must(common.RegisterConfig((*ClientConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
return NewClient(ctx, config.(*ClientConfig))
}))
}
+136
View File
@@ -0,0 +1,136 @@
// Code generated by protoc-gen-go. DO NOT EDIT.
// versions:
// protoc-gen-go v1.36.11
// protoc v6.33.5
// source: proxy/masque/config.proto
package masque
import (
protocol "github.com/xtls/xray-core/common/protocol"
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
reflect "reflect"
sync "sync"
unsafe "unsafe"
)
const (
// Verify that this generated code is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion)
// Verify that runtime/protoimpl is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
)
type ClientConfig struct {
state protoimpl.MessageState `protogen:"open.v1"`
Server *protocol.ServerEndpoint `protobuf:"bytes,1,opt,name=server,proto3" json:"server,omitempty"`
RemoteDns []string `protobuf:"bytes,2,rep,name=remote_dns,json=remoteDns,proto3" json:"remote_dns,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *ClientConfig) Reset() {
*x = ClientConfig{}
mi := &file_proxy_masque_config_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *ClientConfig) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*ClientConfig) ProtoMessage() {}
func (x *ClientConfig) ProtoReflect() protoreflect.Message {
mi := &file_proxy_masque_config_proto_msgTypes[0]
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
}
return ms
}
return mi.MessageOf(x)
}
// Deprecated: Use ClientConfig.ProtoReflect.Descriptor instead.
func (*ClientConfig) Descriptor() ([]byte, []int) {
return file_proxy_masque_config_proto_rawDescGZIP(), []int{0}
}
func (x *ClientConfig) GetServer() *protocol.ServerEndpoint {
if x != nil {
return x.Server
}
return nil
}
func (x *ClientConfig) GetRemoteDns() []string {
if x != nil {
return x.RemoteDns
}
return nil
}
var File_proxy_masque_config_proto protoreflect.FileDescriptor
const file_proxy_masque_config_proto_rawDesc = "" +
"\n" +
"\x19proxy/masque/config.proto\x12\x11xray.proxy.masque\x1a!common/protocol/server_spec.proto\"k\n" +
"\fClientConfig\x12<\n" +
"\x06server\x18\x01 \x01(\v2$.xray.common.protocol.ServerEndpointR\x06server\x12\x1d\n" +
"\n" +
"remote_dns\x18\x02 \x03(\tR\tremoteDnsBU\n" +
"\x15com.xray.proxy.masqueP\x01Z&github.com/xtls/xray-core/proxy/masque\xaa\x02\x11Xray.Proxy.Masqueb\x06proto3"
var (
file_proxy_masque_config_proto_rawDescOnce sync.Once
file_proxy_masque_config_proto_rawDescData []byte
)
func file_proxy_masque_config_proto_rawDescGZIP() []byte {
file_proxy_masque_config_proto_rawDescOnce.Do(func() {
file_proxy_masque_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_proxy_masque_config_proto_rawDesc), len(file_proxy_masque_config_proto_rawDesc)))
})
return file_proxy_masque_config_proto_rawDescData
}
var file_proxy_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
var file_proxy_masque_config_proto_goTypes = []any{
(*ClientConfig)(nil), // 0: xray.proxy.masque.ClientConfig
(*protocol.ServerEndpoint)(nil), // 1: xray.common.protocol.ServerEndpoint
}
var file_proxy_masque_config_proto_depIdxs = []int32{
1, // 0: xray.proxy.masque.ClientConfig.server:type_name -> xray.common.protocol.ServerEndpoint
1, // [1:1] is the sub-list for method output_type
1, // [1:1] is the sub-list for method input_type
1, // [1:1] is the sub-list for extension type_name
1, // [1:1] is the sub-list for extension extendee
0, // [0:1] is the sub-list for field type_name
}
func init() { file_proxy_masque_config_proto_init() }
func file_proxy_masque_config_proto_init() {
if File_proxy_masque_config_proto != nil {
return
}
type x struct{}
out := protoimpl.TypeBuilder{
File: protoimpl.DescBuilder{
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_masque_config_proto_rawDesc), len(file_proxy_masque_config_proto_rawDesc)),
NumEnums: 0,
NumMessages: 1,
NumExtensions: 0,
NumServices: 0,
},
GoTypes: file_proxy_masque_config_proto_goTypes,
DependencyIndexes: file_proxy_masque_config_proto_depIdxs,
MessageInfos: file_proxy_masque_config_proto_msgTypes,
}.Build()
File_proxy_masque_config_proto = out.File
file_proxy_masque_config_proto_goTypes = nil
file_proxy_masque_config_proto_depIdxs = nil
}
+14
View File
@@ -0,0 +1,14 @@
syntax = "proto3";
package xray.proxy.masque;
option csharp_namespace = "Xray.Proxy.Masque";
option go_package = "github.com/xtls/xray-core/proxy/masque";
option java_package = "com.xray.proxy.masque";
option java_multiple_files = true;
import "common/protocol/server_spec.proto";
message ClientConfig {
xray.common.protocol.ServerEndpoint server = 1;
repeated string remote_dns = 2;
}
+20 -10
View File
@@ -199,9 +199,9 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
return errors.New("failed to create UDP connection").Base(err)
}
defer conn.Close()
c := &udpConnClient{
c := &UDPConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
dest: conn.RemoteAddr().(*net.UDPAddr),
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = c
writer = c
@@ -336,6 +336,9 @@ func (h *Handler) init(ctx context.Context) error {
}
func (h *Handler) resolveLocal(host string) (net.IP, error) {
if ip := net.ParseIP(host); ip != nil {
return ip, nil
}
ips, _, err := h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
if err != nil {
return nil, err
@@ -375,12 +378,12 @@ func (h *Handler) resolveLocal(host string) (net.IP, error) {
return got[dice.Roll(len(got))], nil
}
type udpConnClient struct {
type UDPConnClient struct {
net.PacketConn
dest *net.UDPAddr
Dest *net.UDPAddr
}
func (c *udpConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
func (c *UDPConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
b := buf.New()
b.Resize(0, buf.Size)
n, addr, err := c.PacketConn.ReadFrom(b.Bytes())
@@ -399,9 +402,9 @@ func (c *udpConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
return buf.MultiBuffer{b}, nil
}
func (c *udpConnClient) WriteMultiBuffer(mb buf.MultiBuffer) error {
func (c *UDPConnClient) WriteMultiBuffer(mb buf.MultiBuffer) error {
for i, b := range mb {
dst := c.dest
dst := c.Dest
if b.UDP != nil {
if b.UDP.Address.Family().IsDomain() {
if b.UDP.Port != net.Port(dst.Port) {
@@ -459,20 +462,27 @@ func (c *cache) run() {
return
}
c.running = true
c.m = make(map[string]entry)
if c.m == nil {
c.m = make(map[string]entry)
}
go c.gc()
}
func (c *cache) gc() {
ticker := time.NewTicker(time.Minute)
for {
now := <-ticker.C
defer ticker.Stop()
for now := range ticker.C {
c.mu.Lock()
for key, entry := range c.m {
if now.After(entry.deadline) {
delete(c.m, key)
}
}
if len(c.m) == 0 {
c.running = false
c.mu.Unlock()
return
}
c.mu.Unlock()
}
}
+275
View File
@@ -0,0 +1,275 @@
package scenarios
import (
"context"
gotls "crypto/tls"
"crypto/x509"
go_errors "errors"
"io"
"net/http"
"net/netip"
"sync/atomic"
"testing"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"golang.org/x/sync/errgroup"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
"gvisor.dev/gvisor/pkg/tcpip/network/ipv4"
"gvisor.dev/gvisor/pkg/tcpip/network/ipv6"
"github.com/xtls/xray-core/app/log"
"github.com/xtls/xray-core/app/proxyman"
"github.com/xtls/xray-core/common"
clog "github.com/xtls/xray-core/common/log"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/protocol/tls/cert"
"github.com/xtls/xray-core/common/serial"
core "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/proxy/dokodemo"
"github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/proxy/wireguard"
"github.com/xtls/xray-core/testing/servers/tcp"
"github.com/xtls/xray-core/testing/servers/udp"
"github.com/xtls/xray-core/transport/internet"
transmasque "github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/masque/connectip"
"github.com/xtls/xray-core/transport/internet/tls"
)
var (
masqueServerV4 = netip.MustParseAddr("10.13.0.1")
masqueServerV6 = netip.MustParseAddr("fd13::1")
masqueClientV4 = netip.MustParsePrefix("10.13.0.2/32")
masqueClientV6 = netip.MustParsePrefix("fd13::2/128")
)
const (
masqueEchoPort = 7
masqueAuthorization = "Basic dTpw"
)
func startMasqueServer(t *testing.T) (net.Port, [32]byte) {
dev, _, gstack, err := wireguard.CreateNetTUN([]netip.Addr{masqueServerV4, masqueServerV6}, nil, transmasque.MinPacketSize, false)
common.Must(err)
t.Cleanup(func() { dev.Close() })
for _, addr := range []netip.Addr{masqueServerV4, masqueServerV6} {
proto := ipv4.ProtocolNumber
if addr.Is6() {
proto = ipv6.ProtocolNumber
}
local := tcpip.FullAddress{NIC: 1, Addr: tcpip.AddrFromSlice(addr.AsSlice()), Port: masqueEchoPort}
l, err := gonet.ListenTCP(gstack, local, proto)
common.Must(err)
go func() {
for {
c, err := l.Accept()
if err != nil {
return
}
go func() {
defer c.Close()
b := make([]byte, 2048)
for {
n, err := c.Read(b)
if err != nil {
return
}
if _, err := c.Write(xor(b[:n])); err != nil {
return
}
}
}()
}
}()
u, err := gonet.DialUDP(gstack, &local, nil, proto)
common.Must(err)
go func() {
b := make([]byte, 2048)
for {
n, addr, err := u.ReadFrom(b)
if err != nil {
return
}
u.WriteTo(xor(b[:n]), addr)
}
}()
}
var current atomic.Pointer[connectip.Conn]
go func() {
bufs := [][]byte{make([]byte, transmasque.MinPacketSize)}
sizes := []int{0}
for {
if _, err := dev.Read(bufs, sizes, 0); err != nil {
return
}
if conn := current.Load(); conn != nil {
if icmp, _ := conn.WritePacket(bufs[0][:sizes[0]]); len(icmp) > 0 {
go dev.Write([][]byte{icmp}, 0)
}
}
}
}()
handler := func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != transmasque.DefaultPath {
w.WriteHeader(http.StatusNotFound)
return
}
if r.Header.Get("Authorization") != masqueAuthorization {
w.WriteHeader(http.StatusUnauthorized)
return
}
req, err := connectip.ParseProxyRequest(r)
if err != nil {
var perr *connectip.ProxyRequestParseError
if go_errors.As(err, &perr) {
w.WriteHeader(perr.HTTPStatus)
}
return
}
conn, err := (&connectip.Proxy{}).Proxy(w, req)
if err != nil {
return
}
defer conn.Close()
common.Must(conn.AssignAddresses([]netip.Prefix{masqueClientV4, masqueClientV6}))
common.Must(conn.AdvertiseRoute([]connectip.IPRoute{
{StartIP: netip.IPv4Unspecified(), EndIP: netip.AddrFrom4([4]byte{255, 255, 255, 255})},
{StartIP: netip.IPv6Unspecified(), EndIP: netip.AddrFrom16([16]byte{0: 0xff, 1: 0xff, 2: 0xff, 3: 0xff, 4: 0xff, 5: 0xff, 6: 0xff, 7: 0xff, 8: 0xff, 9: 0xff, 10: 0xff, 11: 0xff, 12: 0xff, 13: 0xff, 14: 0xff, 15: 0xff})},
}))
go func() {
for {
ar, err := conn.ReceiveAddressRequest(context.Background())
if err != nil {
return
}
assigned := make([]netip.Prefix, len(ar.Prefixes))
for i, p := range ar.Prefixes {
if p.Addr().Is4() {
assigned[i] = masqueClientV4
} else {
assigned[i] = masqueClientV6
}
}
ar.Respond(assigned, nil)
}
}()
current.Store(conn)
b := make([]byte, 2048)
for {
n, err := conn.ReadPacket(b)
if err != nil {
if go_errors.Is(err, io.ErrShortBuffer) {
continue
}
return
}
dev.Write([][]byte{b[:n]}, 0)
}
}
certificate, certHash := cert.MustGenerate(nil, cert.CommonName("localhost"))
key := common.Must2(x509.ParsePKCS8PrivateKey(certificate.PrivateKey))
tlsConfig := &gotls.Config{
Certificates: []gotls.Certificate{{Certificate: [][]byte{certificate.Certificate}, PrivateKey: key}},
NextProtos: []string{http3.NextProtoH3},
}
pktConn := common.Must2(net.ListenUDP("udp", &net.UDPAddr{IP: net.LocalHostIP.IP()}))
tr := &quic.Transport{Conn: pktConn}
ln := common.Must2(tr.ListenEarly(tlsConfig, &quic.Config{EnableDatagrams: true, InitialPacketSize: 1350}))
server := &http3.Server{Handler: http.HandlerFunc(handler), EnableDatagrams: true}
go server.ServeListener(ln)
t.Cleanup(func() {
server.Close()
ln.Close()
tr.Close()
pktConn.Close()
})
return net.Port(pktConn.LocalAddr().(*net.UDPAddr).Port), certHash
}
func TestMasque(t *testing.T) {
serverPort, certHash := startMasqueServer(t)
tcpPort := tcp.PickPort()
tcp6Port := tcp.PickPort()
udpPort := udp.PickPort()
dokodemoTo := func(port net.Port, addr netip.Addr, network net.Network) *core.InboundHandlerConfig {
return &core.InboundHandlerConfig{
ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{
PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(port)}},
Listen: net.NewIPOrDomain(net.LocalHostIP),
}),
ProxySettings: serial.ToTypedMessage(&dokodemo.Config{
RewriteAddress: net.NewIPOrDomain(net.IPAddress(addr.AsSlice())),
RewritePort: masqueEchoPort,
AllowedNetworks: []net.Network{network},
}),
}
}
clientConfig := &core.Config{
App: []*serial.TypedMessage{
serial.ToTypedMessage(&log.Config{
ErrorLogLevel: clog.Severity_Debug,
ErrorLogType: log.LogType_Console,
}),
},
Inbound: []*core.InboundHandlerConfig{
dokodemoTo(tcpPort, masqueServerV4, net.Network_TCP),
dokodemoTo(tcp6Port, masqueServerV6, net.Network_TCP),
dokodemoTo(udpPort, masqueServerV4, net.Network_UDP),
},
Outbound: []*core.OutboundHandlerConfig{
{
ProxySettings: serial.ToTypedMessage(&masque.ClientConfig{
Server: &protocol.ServerEndpoint{
Address: net.NewIPOrDomain(net.LocalHostIP),
Port: uint32(serverPort),
},
}),
SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{
StreamSettings: &internet.StreamConfig{
ProtocolName: "masque",
TransportSettings: []*internet.TransportConfig{
{
ProtocolName: "masque",
Settings: serial.ToTypedMessage(&transmasque.Config{
Path: transmasque.DefaultPath,
Headers: map[string]string{"Authorization": masqueAuthorization},
}),
},
},
SecurityType: serial.GetMessageType(&tls.Config{}),
SecuritySettings: []*serial.TypedMessage{
serial.ToTypedMessage(&tls.Config{
ServerName: "localhost",
PinnedPeerCertSha256: [][]byte{certHash[:]},
}),
},
},
}),
},
},
}
servers, err := InitializeServerConfigs(clientConfig)
common.Must(err)
defer CloseAllServers(servers)
var errg errgroup.Group
for range 3 {
errg.Go(testTCPConn(tcpPort, 1024*1024, time.Second*20))
}
errg.Go(testTCPConn(tcp6Port, 1024*1024, time.Second*20))
errg.Go(testUDPConn(udpPort, 1024, time.Second*5))
if err := errg.Wait(); err != nil {
t.Error(err)
}
}
+18
View File
@@ -0,0 +1,18 @@
package masque
import (
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/transport/internet"
)
const protocolName = "masque"
const DefaultPath = "/.well-known/masque/ip/*/*/"
func init() {
common.Must(internet.RegisterProtocolConfigCreator(protocolName, func() interface{} {
return &Config{
Path: DefaultPath,
}
}))
}
+146
View File
@@ -0,0 +1,146 @@
// Code generated by protoc-gen-go. DO NOT EDIT.
// versions:
// protoc-gen-go v1.36.11
// protoc v6.33.5
// source: transport/internet/masque/config.proto
package masque
import (
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
reflect "reflect"
sync "sync"
unsafe "unsafe"
)
const (
// Verify that this generated code is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion)
// Verify that runtime/protoimpl is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
)
type Config struct {
state protoimpl.MessageState `protogen:"open.v1"`
Host string `protobuf:"bytes,1,opt,name=host,proto3" json:"host,omitempty"`
Path string `protobuf:"bytes,2,opt,name=path,proto3" json:"path,omitempty"`
Headers map[string]string `protobuf:"bytes,3,rep,name=headers,proto3" json:"headers,omitempty" protobuf_key:"bytes,1,opt,name=key" protobuf_val:"bytes,2,opt,name=value"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Config) Reset() {
*x = Config{}
mi := &file_transport_internet_masque_config_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *Config) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*Config) ProtoMessage() {}
func (x *Config) ProtoReflect() protoreflect.Message {
mi := &file_transport_internet_masque_config_proto_msgTypes[0]
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
}
return ms
}
return mi.MessageOf(x)
}
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
func (*Config) Descriptor() ([]byte, []int) {
return file_transport_internet_masque_config_proto_rawDescGZIP(), []int{0}
}
func (x *Config) GetHost() string {
if x != nil {
return x.Host
}
return ""
}
func (x *Config) GetPath() string {
if x != nil {
return x.Path
}
return ""
}
func (x *Config) GetHeaders() map[string]string {
if x != nil {
return x.Headers
}
return nil
}
var File_transport_internet_masque_config_proto protoreflect.FileDescriptor
const file_transport_internet_masque_config_proto_rawDesc = "" +
"\n" +
"&transport/internet/masque/config.proto\x12\x1exray.transport.internet.masque\"\xbb\x01\n" +
"\x06Config\x12\x12\n" +
"\x04host\x18\x01 \x01(\tR\x04host\x12\x12\n" +
"\x04path\x18\x02 \x01(\tR\x04path\x12M\n" +
"\aheaders\x18\x03 \x03(\v23.xray.transport.internet.masque.Config.HeadersEntryR\aheaders\x1a:\n" +
"\fHeadersEntry\x12\x10\n" +
"\x03key\x18\x01 \x01(\tR\x03key\x12\x14\n" +
"\x05value\x18\x02 \x01(\tR\x05value:\x028\x01B|\n" +
"\"com.xray.transport.internet.masqueP\x01Z3github.com/xtls/xray-core/transport/internet/masque\xaa\x02\x1eXray.Transport.Internet.Masqueb\x06proto3"
var (
file_transport_internet_masque_config_proto_rawDescOnce sync.Once
file_transport_internet_masque_config_proto_rawDescData []byte
)
func file_transport_internet_masque_config_proto_rawDescGZIP() []byte {
file_transport_internet_masque_config_proto_rawDescOnce.Do(func() {
file_transport_internet_masque_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_transport_internet_masque_config_proto_rawDesc), len(file_transport_internet_masque_config_proto_rawDesc)))
})
return file_transport_internet_masque_config_proto_rawDescData
}
var file_transport_internet_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2)
var file_transport_internet_masque_config_proto_goTypes = []any{
(*Config)(nil), // 0: xray.transport.internet.masque.Config
nil, // 1: xray.transport.internet.masque.Config.HeadersEntry
}
var file_transport_internet_masque_config_proto_depIdxs = []int32{
1, // 0: xray.transport.internet.masque.Config.headers:type_name -> xray.transport.internet.masque.Config.HeadersEntry
1, // [1:1] is the sub-list for method output_type
1, // [1:1] is the sub-list for method input_type
1, // [1:1] is the sub-list for extension type_name
1, // [1:1] is the sub-list for extension extendee
0, // [0:1] is the sub-list for field type_name
}
func init() { file_transport_internet_masque_config_proto_init() }
func file_transport_internet_masque_config_proto_init() {
if File_transport_internet_masque_config_proto != nil {
return
}
type x struct{}
out := protoimpl.TypeBuilder{
File: protoimpl.DescBuilder{
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_masque_config_proto_rawDesc), len(file_transport_internet_masque_config_proto_rawDesc)),
NumEnums: 0,
NumMessages: 2,
NumExtensions: 0,
NumServices: 0,
},
GoTypes: file_transport_internet_masque_config_proto_goTypes,
DependencyIndexes: file_transport_internet_masque_config_proto_depIdxs,
MessageInfos: file_transport_internet_masque_config_proto_msgTypes,
}.Build()
File_transport_internet_masque_config_proto = out.File
file_transport_internet_masque_config_proto_goTypes = nil
file_transport_internet_masque_config_proto_depIdxs = nil
}
+13
View File
@@ -0,0 +1,13 @@
syntax = "proto3";
package xray.transport.internet.masque;
option csharp_namespace = "Xray.Transport.Internet.Masque";
option go_package = "github.com/xtls/xray-core/transport/internet/masque";
option java_package = "com.xray.transport.internet.masque";
option java_multiple_files = true;
message Config {
string host = 1;
string path = 2;
map<string, string> headers = 3;
}
+115
View File
@@ -0,0 +1,115 @@
package masque
import (
"context"
go_errors "errors"
"net/netip"
"slices"
"sync"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/masque/connectip"
)
type PacketTooBigError struct {
ICMP []byte
}
func (e *PacketTooBigError) Error() string {
return "packet too big for the tunnel"
}
type Conn struct {
ipConn *connectip.Conn
quicConn *quic.Conn
local []netip.Addr
closeOnce sync.Once
}
func (c *Conn) LocalAddrs() []netip.Addr {
return c.local
}
func (c *Conn) Read(b []byte) (int, error) {
return c.ipConn.ReadPacket(b)
}
func (c *Conn) Write(b []byte) (int, error) {
icmp, err := c.ipConn.WritePacket(b)
if err != nil {
if go_errors.Is(err, connectip.ErrMTUTooSmall) {
errors.LogWarning(context.Background(), "MASQUE: closing the tunnel as it cannot carry ", MinPacketSize, "-byte packets")
} else {
errors.LogInfoInner(context.Background(), err, "MASQUE: closing the tunnel as sending failed")
}
c.Close()
return 0, err
}
if len(icmp) > 0 {
return 0, &PacketTooBigError{ICMP: icmp}
}
return len(b), nil
}
func (c *Conn) Close() error {
c.closeOnce.Do(func() {
c.ipConn.Close()
c.quicConn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "")
})
return nil
}
func (c *Conn) LocalAddr() net.Addr {
return c.quicConn.LocalAddr()
}
func (c *Conn) RemoteAddr() net.Addr {
return c.quicConn.RemoteAddr()
}
func (c *Conn) SetDeadline(time.Time) error {
return nil
}
func (c *Conn) SetReadDeadline(time.Time) error {
return nil
}
func (c *Conn) SetWriteDeadline(time.Time) error {
return nil
}
func (c *Conn) serveAddressAssignments() {
for {
assigned, err := c.ipConn.ReceiveAddressAssignment(context.Background())
if err != nil {
return
}
for _, addr := range c.local {
if !slices.ContainsFunc(assigned, func(a connectip.AssignedAddress) bool { return !a.Rejected() && a.IPPrefix.Contains(addr) }) {
errors.LogInfo(context.Background(), "MASQUE: closing the tunnel as the proxy withdrew ", addr)
c.Close()
return
}
}
if len(localAddrs(assigned)) > len(c.local) {
errors.LogInfo(context.Background(), "MASQUE: the proxy assigned another IP family, which is used once the tunnel is set up again")
}
}
}
func (c *Conn) serveAddressRequests() {
for {
req, err := c.ipConn.ReceiveAddressRequest(context.Background())
if err != nil {
return
}
if err := req.Respond(make([]netip.Prefix, len(req.Prefixes)), nil); err != nil {
return
}
}
}
@@ -0,0 +1,7 @@
Copyright 2024 Marten Seemann
Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
@@ -0,0 +1,86 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"errors"
"fmt"
"net/netip"
"slices"
"sync/atomic"
)
var (
rejectedIPv4Prefix = netip.PrefixFrom(netip.IPv4Unspecified(), 32)
rejectedIPv6Prefix = netip.PrefixFrom(netip.IPv6Unspecified(), 128)
)
type AddressRequestID uint64
type AddressRequest struct {
Prefixes []netip.Prefix
conn *Conn
requested *addressRequestCapsule
responded *atomic.Bool
}
func newAddressRequest(conn *Conn, requested *addressRequestCapsule) *AddressRequest {
return &AddressRequest{
Prefixes: slices.Clone(requested.Prefixes),
conn: conn,
requested: requested,
responded: &atomic.Bool{},
}
}
func (r *AddressRequest) Respond(assignments, additional []netip.Prefix) error {
if r.conn == nil {
return errors.New("connect-ip: invalid address request")
}
if len(assignments) != len(r.requested.RequestIDs) {
return fmt.Errorf(
"connect-ip: expected %d address assignments, got %d",
len(r.requested.RequestIDs),
len(assignments),
)
}
capsule := &addressAssignCapsule{
AssignedAddresses: make([]AssignedAddress, 0, len(assignments)+len(additional)),
}
var zeroPrefix netip.Prefix
for i, p := range assignments {
if p == zeroPrefix {
if r.requested.Prefixes[i].Addr().Is4() {
p = rejectedIPv4Prefix
} else {
p = rejectedIPv6Prefix
}
} else if !p.IsValid() || p != p.Masked() {
return fmt.Errorf("connect-ip: invalid assigned prefix %d: %s", i, p)
}
capsule.AssignedAddresses = append(
capsule.AssignedAddresses,
AssignedAddress{RequestID: r.requested.RequestIDs[i], IPPrefix: p},
)
}
for i, p := range additional {
if !p.IsValid() || p != p.Masked() {
return fmt.Errorf("connect-ip: invalid additional prefix %d: %s", i, p)
}
capsule.AssignedAddresses = append(capsule.AssignedAddresses, AssignedAddress{IPPrefix: p})
}
if !r.responded.CompareAndSwap(false, true) {
return errors.New("connect-ip: address request already answered")
}
restrictPeer := slices.ContainsFunc(capsule.AssignedAddresses, func(a AssignedAddress) bool { return !a.Rejected() })
if err := r.conn.sendAddressAssignment(capsule, restrictPeer); err != nil {
r.responded.Store(false)
return err
}
return nil
}
@@ -0,0 +1,85 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"context"
"net/netip"
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestAddressRequests(t *testing.T) {
client, server := setupConns(t)
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
prefixes := []netip.Prefix{
netip.MustParsePrefix("0.0.0.0/32"),
netip.MustParsePrefix("0.0.0.0/32"),
netip.MustParsePrefix("::/64"),
}
ids, err := client.RequestAddresses(prefixes)
require.NoError(t, err)
require.Equal(t, []AddressRequestID{1, 2, 3}, ids)
req, err := server.ReceiveAddressRequest(ctx)
require.NoError(t, err)
require.Equal(t, prefixes, req.Prefixes)
assignments := []netip.Prefix{netip.MustParsePrefix("192.0.2.1/32"), {}, {}}
additional := []netip.Prefix{netip.MustParsePrefix("2001:db8::/64")}
require.NoError(t, req.Respond(assignments, additional))
received, err := client.ReceiveAddressAssignment(ctx)
require.NoError(t, err)
require.Len(t, received, 4)
require.Equal(t, AssignedAddress{RequestID: ids[0], IPPrefix: assignments[0]}, received[0])
require.Equal(t, ids[1], received[1].RequestID)
require.True(t, received[1].Rejected())
require.Equal(t, ids[2], received[2].RequestID)
require.True(t, received[2].Rejected())
require.Equal(t, AssignedAddress{IPPrefix: additional[0]}, received[3])
ids, err = client.RequestAddresses(prefixes[:1])
require.NoError(t, err)
require.Equal(t, []AddressRequestID{4}, ids)
}
func TestAddressRequestValidation(t *testing.T) {
conn := newProxiedConn(&mockStream{})
defer conn.Close()
for _, prefixes := range [][]netip.Prefix{
nil,
{{}},
{netip.MustParsePrefix("192.0.2.1/24")},
{netip.MustParsePrefix("2001:db8::1/64")},
} {
ids, err := conn.RequestAddresses(prefixes)
require.Error(t, err)
require.Nil(t, ids)
}
}
func TestAddressResponseValidation(t *testing.T) {
conn := newProxiedConn(&mockStream{})
defer conn.Close()
prefixes := []netip.Prefix{netip.MustParsePrefix("192.0.2.1/32")}
req := newAddressRequest(conn, &addressRequestCapsule{RequestIDs: []AddressRequestID{1}, Prefixes: prefixes})
require.ErrorContains(t, (&AddressRequest{}).Respond(nil, nil), "invalid address request")
require.ErrorContains(t, req.Respond(nil, nil), "expected 1 address assignments")
require.ErrorContains(t, req.Respond(prefixes, []netip.Prefix{{}}), "invalid additional prefix")
require.ErrorContains(t,
req.Respond([]netip.Prefix{netip.MustParsePrefix("192.0.2.1/24")}, nil),
"invalid assigned prefix",
)
copied := *req
require.NoError(t, req.Respond(prefixes, nil))
require.ErrorContains(t, copied.Respond(prefixes, nil), "already answered")
}
@@ -0,0 +1,308 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"cmp"
"encoding/binary"
"errors"
"fmt"
"io"
"net/netip"
"github.com/apernet/quic-go/http3"
"github.com/apernet/quic-go/quicvarint"
)
const (
capsuleTypeDatagram http3.CapsuleType = 0
capsuleTypeAddressAssign http3.CapsuleType = 1
capsuleTypeAddressRequest http3.CapsuleType = 2
capsuleTypeRouteAdvertisement http3.CapsuleType = 3
)
const (
maxAddressesPerCapsule = 8192
maxRoutesPerCapsule = 8192
)
type addressAssignCapsule struct {
AssignedAddresses []AssignedAddress
}
type AssignedAddress struct {
RequestID AddressRequestID
IPPrefix netip.Prefix
}
func (a AssignedAddress) Rejected() bool {
return a.IPPrefix == rejectedIPv4Prefix || a.IPPrefix == rejectedIPv6Prefix
}
func (a AssignedAddress) len() int {
return quicvarint.Len(uint64(a.RequestID)) + 1 + a.IPPrefix.Addr().BitLen()/8 + 1
}
type addressRequestCapsule struct {
RequestIDs []AddressRequestID
Prefixes []netip.Prefix
}
func parseAddressAssignCapsule(r http3.CapsuleReader) (*addressAssignCapsule, error) {
var assignedAddresses []AssignedAddress
for r.Remaining() > 0 {
if len(assignedAddresses) >= maxAddressesPerCapsule {
return nil, fmt.Errorf("%w: ADDRESS_ASSIGN capsule contains too many addresses (maximum %d)", errCapsuleLimit, maxAddressesPerCapsule)
}
requestID, prefix, err := parseAddress(r)
if err != nil {
return nil, err
}
assignedAddresses = append(assignedAddresses, AssignedAddress{RequestID: AddressRequestID(requestID), IPPrefix: prefix})
}
return &addressAssignCapsule{AssignedAddresses: assignedAddresses}, nil
}
func (c *addressAssignCapsule) append(b []byte) []byte {
totalLen := 0
for _, addr := range c.AssignedAddresses {
totalLen += addr.len()
}
b = quicvarint.Append(b, uint64(capsuleTypeAddressAssign))
b = quicvarint.Append(b, uint64(totalLen))
for _, addr := range c.AssignedAddresses {
b = quicvarint.Append(b, uint64(addr.RequestID))
if addr.IPPrefix.Addr().Is4() {
b = append(b, 4)
} else {
b = append(b, 6)
}
b = append(b, addr.IPPrefix.Addr().AsSlice()...)
b = append(b, byte(addr.IPPrefix.Bits()))
}
return b
}
func parseAddressRequestCapsule(r http3.CapsuleReader) (*addressRequestCapsule, error) {
if r.Remaining() == 0 {
return nil, errors.New("ADDRESS_REQUEST capsule contains no addresses")
}
capsule := &addressRequestCapsule{}
for r.Remaining() > 0 {
if len(capsule.Prefixes) >= maxAddressesPerCapsule {
return nil, fmt.Errorf("%w: ADDRESS_REQUEST capsule contains too many addresses (maximum %d)", errCapsuleLimit, maxAddressesPerCapsule)
}
requestID, prefix, err := parseAddress(r)
if err != nil {
return nil, err
}
if requestID == 0 {
return nil, errors.New("ADDRESS_REQUEST capsule contains a zero request ID")
}
capsule.RequestIDs = append(capsule.RequestIDs, AddressRequestID(requestID))
capsule.Prefixes = append(capsule.Prefixes, prefix)
}
return capsule, nil
}
func (c *addressRequestCapsule) append(b []byte) []byte {
var totalLen int
for i, p := range c.Prefixes {
totalLen += quicvarint.Len(uint64(c.RequestIDs[i])) + 1 + p.Addr().BitLen()/8 + 1
}
b = quicvarint.Append(b, uint64(capsuleTypeAddressRequest))
b = quicvarint.Append(b, uint64(totalLen))
for i, p := range c.Prefixes {
b = quicvarint.Append(b, uint64(c.RequestIDs[i]))
if p.Addr().Is4() {
b = append(b, 4)
} else {
b = append(b, 6)
}
b = append(b, p.Addr().AsSlice()...)
b = append(b, byte(p.Bits()))
}
return b
}
func parseAddress(r io.Reader) (requestID uint64, prefix netip.Prefix, _ error) {
vr := quicvarint.NewReader(r)
requestID, err := quicvarint.Read(vr)
if err != nil {
return 0, netip.Prefix{}, err
}
ipVersion, err := vr.ReadByte()
if err != nil {
return 0, netip.Prefix{}, err
}
var ip netip.Addr
switch ipVersion {
case 4:
var ipv4 [4]byte
if _, err := io.ReadFull(r, ipv4[:]); err != nil {
return 0, netip.Prefix{}, err
}
ip = netip.AddrFrom4(ipv4)
case 6:
var ipv6 [16]byte
if _, err := io.ReadFull(r, ipv6[:]); err != nil {
return 0, netip.Prefix{}, err
}
ip = netip.AddrFrom16(ipv6)
default:
return 0, netip.Prefix{}, fmt.Errorf("invalid IP version: %d", ipVersion)
}
prefixLen, err := vr.ReadByte()
if err != nil {
return 0, netip.Prefix{}, err
}
if int(prefixLen) > ip.BitLen() {
return 0, netip.Prefix{}, fmt.Errorf("prefix length %d exceeds IP address length (%d)", prefixLen, ip.BitLen())
}
prefix = netip.PrefixFrom(ip, int(prefixLen))
if prefix != prefix.Masked() {
return 0, netip.Prefix{}, errors.New("lower bits not covered by prefix length are not all zero")
}
return requestID, prefix, nil
}
type routeAdvertisementCapsule struct {
IPAddressRanges []IPRoute
}
type IPRoute struct {
StartIP netip.Addr
EndIP netip.Addr
IPProtocol uint8
}
func (r IPRoute) len() int { return 1 + r.StartIP.BitLen()/8 + r.EndIP.BitLen()/8 + 1 }
func (r IPRoute) Prefixes() []netip.Prefix { return rangeToPrefixes(r.StartIP, r.EndIP) }
func parseRouteAdvertisementCapsule(r http3.CapsuleReader) (*routeAdvertisementCapsule, error) {
var ranges []IPRoute
for r.Remaining() > 0 {
if len(ranges) >= maxRoutesPerCapsule {
return nil, fmt.Errorf("%w: ROUTE_ADVERTISEMENT capsule contains too many routes (maximum %d)", errCapsuleLimit, maxRoutesPerCapsule)
}
ipRange, err := parseIPAddressRange(r)
if err != nil {
return nil, err
}
if len(ranges) > 0 {
if err := checkRouteOrder(ranges[len(ranges)-1], ipRange); err != nil {
return nil, err
}
}
ranges = append(ranges, ipRange)
}
return &routeAdvertisementCapsule{IPAddressRanges: ranges}, nil
}
func (r IPRoute) validate() error {
if !r.StartIP.IsValid() || !r.EndIP.IsValid() || r.StartIP.Zone() != "" || r.EndIP.Zone() != "" {
return fmt.Errorf("invalid IP address range %s-%s", r.StartIP, r.EndIP)
}
if r.StartIP.Is4() != r.EndIP.Is4() {
return fmt.Errorf("IP address range %s-%s mixes IP versions", r.StartIP, r.EndIP)
}
if r.StartIP.Compare(r.EndIP) > 0 {
return fmt.Errorf("start IP %s is greater than end IP %s", r.StartIP, r.EndIP)
}
return nil
}
func checkRouteOrder(a, b IPRoute) error {
switch cmp.Or(
cmp.Compare(a.StartIP.BitLen(), b.StartIP.BitLen()),
cmp.Compare(a.IPProtocol, b.IPProtocol),
) {
case 1:
return fmt.Errorf("routes are not ordered by IP version and IP protocol: %s-%s (protocol %d) precedes %s-%s (protocol %d)",
a.StartIP, a.EndIP, a.IPProtocol, b.StartIP, b.EndIP, b.IPProtocol)
case 0:
if a.EndIP.Compare(b.StartIP) >= 0 {
return fmt.Errorf("IP address ranges %s-%s and %s-%s (protocol %d) overlap or are not in ascending order",
a.StartIP, a.EndIP, b.StartIP, b.EndIP, b.IPProtocol)
}
}
return nil
}
func (c *routeAdvertisementCapsule) append(b []byte) []byte {
var totalLen int
for _, ipRange := range c.IPAddressRanges {
totalLen += ipRange.len()
}
b = quicvarint.Append(b, uint64(capsuleTypeRouteAdvertisement))
b = quicvarint.Append(b, uint64(totalLen))
for _, ipRange := range c.IPAddressRanges {
if ipRange.StartIP.Is4() {
b = append(b, 4)
} else {
b = append(b, 6)
}
b = append(b, ipRange.StartIP.AsSlice()...)
b = append(b, ipRange.EndIP.AsSlice()...)
b = append(b, ipRange.IPProtocol)
}
return b
}
func parseIPAddressRange(r io.Reader) (IPRoute, error) {
var ipVersion uint8
if err := binary.Read(r, binary.LittleEndian, &ipVersion); err != nil {
return IPRoute{}, err
}
var startIP, endIP netip.Addr
switch ipVersion {
case 4:
var start, end [4]byte
if _, err := io.ReadFull(r, start[:]); err != nil {
return IPRoute{}, err
}
if _, err := io.ReadFull(r, end[:]); err != nil {
return IPRoute{}, err
}
startIP = netip.AddrFrom4(start)
endIP = netip.AddrFrom4(end)
case 6:
var start, end [16]byte
if _, err := io.ReadFull(r, start[:]); err != nil {
return IPRoute{}, err
}
if _, err := io.ReadFull(r, end[:]); err != nil {
return IPRoute{}, err
}
startIP = netip.AddrFrom16(start)
endIP = netip.AddrFrom16(end)
default:
return IPRoute{}, fmt.Errorf("invalid IP version: %d", ipVersion)
}
if startIP.Compare(endIP) > 0 {
return IPRoute{}, errors.New("start IP is greater than end IP")
}
var ipProtocol uint8
if err := binary.Read(r, binary.LittleEndian, &ipProtocol); err != nil {
return IPRoute{}, err
}
return IPRoute{
StartIP: startIP,
EndIP: endIP,
IPProtocol: ipProtocol,
}, nil
}
@@ -0,0 +1,449 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"bytes"
"context"
"io"
"net"
"net/netip"
"testing"
"time"
"github.com/apernet/quic-go/http3"
"github.com/apernet/quic-go/quicvarint"
"github.com/stretchr/testify/require"
)
func newCapsuleReader(t *testing.T, typ http3.CapsuleType, payload []byte) http3.CapsuleReader {
t.Helper()
data := quicvarint.Append(nil, uint64(typ))
data = quicvarint.Append(data, uint64(len(payload)))
data = append(data, payload...)
parsedType, cr, err := http3.NewCapsuleParser(bytes.NewReader(data)).Next()
require.NoError(t, err)
require.Equal(t, typ, parsedType)
return cr
}
func testIncompleteCapsule(t *testing.T, data []byte, parse func(http3.CapsuleReader) error) {
t.Helper()
r := bytes.NewReader(data)
_, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.NoError(t, parse(cr))
require.Zero(t, r.Len())
for i := range data {
_, cr, err := http3.NewCapsuleParser(bytes.NewReader(data[:i])).Next()
if err != nil {
if i == 0 {
require.ErrorIs(t, err, io.EOF)
} else {
require.ErrorIs(t, err, io.ErrUnexpectedEOF)
}
continue
}
require.ErrorIs(t, parse(cr), io.ErrUnexpectedEOF)
}
}
func testCapsuleEntryLimit[T any](t *testing.T, typ http3.CapsuleType, limit int, entry func(i int) []byte, parse func(http3.CapsuleReader) (*T, error)) {
t.Helper()
var payload []byte
for i := range limit {
payload = append(payload, entry(i)...)
}
r := newCapsuleReader(t, typ, payload)
_, err := parse(r)
require.NoError(t, err)
require.Zero(t, r.Remaining())
data := quicvarint.Append(nil, uint64(typ))
data = quicvarint.Append(data, uint64(len(payload)+1))
_, r, err = http3.NewCapsuleParser(bytes.NewReader(append(data, payload...))).Next()
require.NoError(t, err)
_, err = parse(r)
require.ErrorContains(t, err, "too many")
require.Equal(t, int64(1), r.Remaining())
}
func TestParseAddressAssignCapsule(t *testing.T) {
addr1 := quicvarint.Append(nil, 1337)
addr1 = append(addr1, 4)
addr1 = append(addr1, netip.AddrFrom4([4]byte{1, 2, 3, 0}).AsSlice()...)
addr1 = append(addr1, 24)
addr2 := quicvarint.Append(nil, 1338)
addr2 = append(addr2, 6)
addr2 = append(addr2, netip.MustParseAddr("2001:db8::1").AsSlice()...)
addr2 = append(addr2, 128)
data := quicvarint.Append(nil, uint64(capsuleTypeAddressAssign))
data = quicvarint.Append(data, uint64(len(addr1)+len(addr2)))
data = append(data, addr1...)
data = append(data, addr2...)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeAddressAssign, typ)
capsule, err := parseAddressAssignCapsule(cr)
require.NoError(t, err)
require.Equal(t,
[]AssignedAddress{
{RequestID: 1337, IPPrefix: netip.MustParsePrefix("1.2.3.0/24")},
{RequestID: 1338, IPPrefix: netip.MustParsePrefix("2001:db8::1/128")},
},
capsule.AssignedAddresses,
)
require.Zero(t, r.Len())
}
func TestParseAddressAssignCapsuleLimit(t *testing.T) {
entry := []byte{1, 4, 192, 0, 2, 1, 32}
testCapsuleEntryLimit(t, capsuleTypeAddressAssign, maxAddressesPerCapsule, func(int) []byte { return entry }, parseAddressAssignCapsule)
}
func TestAssignedAddressRejected(t *testing.T) {
for _, prefix := range []string{"0.0.0.0/32", "::/128"} {
require.True(t, (AssignedAddress{RequestID: 1, IPPrefix: netip.MustParsePrefix(prefix)}).Rejected())
}
for _, prefix := range []string{"0.0.0.0/0", "0.0.0.0/31", "::/0", "::/127", "192.0.2.1/32", "2001:db8::1/128"} {
require.False(t, (AssignedAddress{RequestID: 1, IPPrefix: netip.MustParsePrefix(prefix)}).Rejected())
}
require.False(t, (AssignedAddress{}).Rejected())
}
func TestWriteAddressAssignCapsule(t *testing.T) {
c := &addressAssignCapsule{
AssignedAddresses: []AssignedAddress{
{RequestID: 1337, IPPrefix: netip.MustParsePrefix("1.2.3.0/24")},
{RequestID: 1338, IPPrefix: netip.MustParsePrefix("2001:db8::1/128")},
},
}
data := c.append(nil)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeAddressAssign, typ)
parsed, err := parseAddressAssignCapsule(cr)
require.NoError(t, err)
require.Equal(t, c, parsed)
require.Zero(t, r.Len())
}
func TestParseAddressAssignCapsuleInvalid(t *testing.T) {
testParseAddressCapsuleInvalid(t, capsuleTypeAddressAssign, func(r http3.CapsuleReader) error {
_, err := parseAddressAssignCapsule(r)
return err
})
}
func testParseAddressCapsuleInvalid(t *testing.T, typ http3.CapsuleType, f func(r http3.CapsuleReader) error) {
t.Run("invalid IP version", func(t *testing.T) {
addr1 := quicvarint.Append(nil, 1337)
addr1 = append(addr1, 5)
addr1 = append(addr1, netip.AddrFrom4([4]byte{1, 2, 3, 4}).AsSlice()...)
addr1 = append(addr1, 32)
require.ErrorContains(t, f(newCapsuleReader(t, typ, addr1)), "invalid IP version: 5")
})
t.Run("invalid prefix length", func(t *testing.T) {
addr1 := quicvarint.Append(nil, 1337)
addr1 = append(addr1, 4)
addr1 = append(addr1, netip.AddrFrom4([4]byte{1, 2, 3, 4}).AsSlice()...)
addr1 = append(addr1, 33)
require.ErrorContains(t, f(newCapsuleReader(t, typ, addr1)), "prefix length 33 exceeds IP address length (32)")
})
t.Run("lower bits not covered by prefix length are not all zero", func(t *testing.T) {
addr1 := quicvarint.Append(nil, 1337)
addr1 = append(addr1, 4)
addr1 = append(addr1, netip.AddrFrom4([4]byte{1, 2, 3, 4}).AsSlice()...)
addr1 = append(addr1, 28)
require.ErrorContains(t, f(newCapsuleReader(t, typ, addr1)), "lower bits not covered by prefix length are not all zero")
})
t.Run("incomplete capsule", func(t *testing.T) {
var data []byte
switch typ {
case capsuleTypeAddressAssign:
data = (&addressAssignCapsule{
AssignedAddresses: []AssignedAddress{
{RequestID: 1337, IPPrefix: netip.MustParsePrefix("1.2.3.4/32")},
{RequestID: 1338, IPPrefix: netip.MustParsePrefix("2001:db8::1/128")},
},
}).append(nil)
case capsuleTypeAddressRequest:
data = (&addressRequestCapsule{
RequestIDs: []AddressRequestID{1337, 1338},
Prefixes: []netip.Prefix{netip.MustParsePrefix("1.2.3.4/32"), netip.MustParsePrefix("2001:db8::1/128")},
}).append(nil)
default:
t.Fatalf("unexpected capsule type: %d", typ)
}
testIncompleteCapsule(t, data, f)
})
}
func TestParseAddressRequestCapsule(t *testing.T) {
addr1 := quicvarint.Append(nil, 1337)
addr1 = append(addr1, 4)
addr1 = append(addr1, netip.AddrFrom4([4]byte{1, 2, 3, 0}).AsSlice()...)
addr1 = append(addr1, 24)
addr2 := quicvarint.Append(nil, 1338)
addr2 = append(addr2, 6)
addr2 = append(addr2, netip.MustParseAddr("2001:db8::1").AsSlice()...)
addr2 = append(addr2, 128)
data := quicvarint.Append(nil, uint64(capsuleTypeAddressRequest))
data = quicvarint.Append(data, uint64(len(addr1)+len(addr2)))
data = append(data, addr1...)
data = append(data, addr2...)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeAddressRequest, typ)
capsule, err := parseAddressRequestCapsule(cr)
require.NoError(t, err)
require.Equal(t, []AddressRequestID{1337, 1338}, capsule.RequestIDs)
require.Equal(t, []netip.Prefix{netip.MustParsePrefix("1.2.3.0/24"), netip.MustParsePrefix("2001:db8::1/128")}, capsule.Prefixes)
require.Zero(t, r.Len())
}
func TestParseAddressRequestCapsuleLimit(t *testing.T) {
entry := []byte{1, 4, 192, 0, 2, 1, 32}
testCapsuleEntryLimit(t, capsuleTypeAddressRequest, maxAddressesPerCapsule, func(int) []byte { return entry }, parseAddressRequestCapsule)
}
func TestWriteAddressRequestCapsule(t *testing.T) {
c := &addressRequestCapsule{
RequestIDs: []AddressRequestID{1337, 1338},
Prefixes: []netip.Prefix{netip.MustParsePrefix("1.2.3.0/24"), netip.MustParsePrefix("2001:db8::1/128")},
}
data := c.append(nil)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeAddressRequest, typ)
parsed, err := parseAddressRequestCapsule(cr)
require.NoError(t, err)
require.Equal(t, c, parsed)
require.Zero(t, r.Len())
}
func TestParseAddressRequestCapsuleInvalid(t *testing.T) {
t.Run("empty", func(t *testing.T) {
_, err := parseAddressRequestCapsule(newCapsuleReader(t, capsuleTypeAddressRequest, nil))
require.ErrorContains(t, err, "contains no addresses")
})
t.Run("zero request ID", func(t *testing.T) {
_, err := parseAddressRequestCapsule(newCapsuleReader(t, capsuleTypeAddressRequest, []byte{0, 4, 192, 0, 2, 1, 32}))
require.ErrorContains(t, err, "zero request ID")
})
testParseAddressCapsuleInvalid(t, capsuleTypeAddressRequest, func(r http3.CapsuleReader) error {
_, err := parseAddressRequestCapsule(r)
return err
})
}
func TestParseRouteAdvertisementCapsule(t *testing.T) {
iprange1 := []byte{4}
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 1, 1, 1}).AsSlice()...)
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 2, 3, 4}).AsSlice()...)
iprange1 = append(iprange1, 13)
iprange2 := []byte{6}
iprange2 = append(iprange2, netip.MustParseAddr("2001:db8::1").AsSlice()...)
iprange2 = append(iprange2, netip.MustParseAddr("2001:db8::100").AsSlice()...)
iprange2 = append(iprange2, 37)
data := quicvarint.Append(nil, uint64(capsuleTypeRouteAdvertisement))
data = quicvarint.Append(data, uint64(len(iprange1)+len(iprange2)))
data = append(data, iprange1...)
data = append(data, iprange2...)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeRouteAdvertisement, typ)
capsule, err := parseRouteAdvertisementCapsule(cr)
require.NoError(t, err)
require.Equal(t,
[]IPRoute{
{StartIP: netip.MustParseAddr("1.1.1.1"), EndIP: netip.MustParseAddr("1.2.3.4"), IPProtocol: 13},
{StartIP: netip.MustParseAddr("2001:db8::1"), EndIP: netip.MustParseAddr("2001:db8::100"), IPProtocol: 37},
},
capsule.IPAddressRanges,
)
require.Equal(t,
rangeToPrefixes(netip.MustParseAddr("1.1.1.1"), netip.MustParseAddr("1.2.3.4")),
capsule.IPAddressRanges[0].Prefixes(),
)
require.Equal(t,
rangeToPrefixes(netip.MustParseAddr("2001:db8::1"), netip.MustParseAddr("2001:db8::100")),
capsule.IPAddressRanges[1].Prefixes(),
)
require.Zero(t, r.Len())
}
func TestParseRouteAdvertisementCapsuleLimit(t *testing.T) {
entry := func(i int) []byte { return []byte{4, 10, 0, byte(i >> 8), byte(i), 10, 0, byte(i >> 8), byte(i), 0} }
testCapsuleEntryLimit(t, capsuleTypeRouteAdvertisement, maxRoutesPerCapsule, entry, parseRouteAdvertisementCapsule)
}
func TestWriteRouteAdvertisementCapsule(t *testing.T) {
c := &routeAdvertisementCapsule{
IPAddressRanges: []IPRoute{
{StartIP: netip.MustParseAddr("1.1.1.1"), EndIP: netip.MustParseAddr("1.2.3.4"), IPProtocol: 13},
{StartIP: netip.MustParseAddr("2001:db8::1"), EndIP: netip.MustParseAddr("2001:db8::100"), IPProtocol: 37},
},
}
data := c.append(nil)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeRouteAdvertisement, typ)
parsed, err := parseRouteAdvertisementCapsule(cr)
require.NoError(t, err)
require.Equal(t, c, parsed)
require.Zero(t, r.Len())
}
func TestParseRouteAdvertisementCapsuleInvalid(t *testing.T) {
t.Run("invalid IP version", func(t *testing.T) {
iprange1 := []byte{5}
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 1, 1, 1}).AsSlice()...)
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 1, 1, 2}).AsSlice()...)
iprange1 = append(iprange1, 13)
_, err := parseRouteAdvertisementCapsule(newCapsuleReader(t, capsuleTypeRouteAdvertisement, iprange1))
require.ErrorContains(t, err, "invalid IP version: 5")
})
t.Run("start IP is greater than end IP", func(t *testing.T) {
iprange1 := []byte{4}
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 2, 3, 4}).AsSlice()...)
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 1, 1, 1}).AsSlice()...)
iprange1 = append(iprange1, 13)
_, err := parseRouteAdvertisementCapsule(newCapsuleReader(t, capsuleTypeRouteAdvertisement, iprange1))
require.ErrorContains(t, err, "start IP is greater than end IP")
})
t.Run("incomplete capsule", func(t *testing.T) {
data := (&routeAdvertisementCapsule{
IPAddressRanges: []IPRoute{
{StartIP: netip.MustParseAddr("1.1.1.1"), EndIP: netip.MustParseAddr("2.2.2.2"), IPProtocol: 13},
{StartIP: netip.MustParseAddr("2001:db8::1"), EndIP: netip.MustParseAddr("2001:db8::100"), IPProtocol: 37},
},
}).append(nil)
testIncompleteCapsule(t, data, func(r http3.CapsuleReader) error {
_, err := parseRouteAdvertisementCapsule(r)
return err
})
})
}
var (
route4a = IPRoute{StartIP: netip.MustParseAddr("10.0.0.0"), EndIP: netip.MustParseAddr("10.0.0.9")}
route4b = IPRoute{StartIP: netip.MustParseAddr("10.0.0.10"), EndIP: netip.MustParseAddr("10.0.0.20")}
route4ab = IPRoute{StartIP: netip.MustParseAddr("10.0.0.9"), EndIP: netip.MustParseAddr("10.0.0.20")}
route6 = IPRoute{StartIP: netip.MustParseAddr("2001:db8::"), EndIP: netip.MustParseAddr("2001:db8::ffff")}
)
func withProtocol(r IPRoute, proto uint8) IPRoute {
r.IPProtocol = proto
return r
}
var routeOrderTests = []struct {
name string
routes []IPRoute
err string
}{
{name: "empty"},
{name: "adjacent ranges", routes: []IPRoute{route4a, route4b}},
{name: "IPv4 before IPv6 with a lower IP protocol", routes: []IPRoute{withProtocol(route4a, 17), route6}},
{name: "same range for different IP protocols", routes: []IPRoute{withProtocol(route4a, 6), withProtocol(route4a, 17)}},
{name: "IP protocol order before address order", routes: []IPRoute{withProtocol(route4b, 6), withProtocol(route4a, 17)}},
{name: "IPv6 before IPv4", routes: []IPRoute{route6, route4a}, err: "not ordered by IP version and IP protocol"},
{name: "descending IP protocols", routes: []IPRoute{withProtocol(route4a, 17), withProtocol(route4b, 6)}, err: "not ordered by IP version and IP protocol"},
{name: "descending ranges", routes: []IPRoute{route4b, route4a}, err: "overlap or are not in ascending order"},
{name: "overlapping ranges", routes: []IPRoute{route4a, route4ab}, err: "overlap or are not in ascending order"},
{name: "duplicate range", routes: []IPRoute{route6, route6}, err: "overlap or are not in ascending order"},
}
func TestParseRouteAdvertisementCapsuleOrder(t *testing.T) {
for _, tc := range routeOrderTests {
t.Run(tc.name, func(t *testing.T) {
data := (&routeAdvertisementCapsule{IPAddressRanges: tc.routes}).append(nil)
_, cr, err := http3.NewCapsuleParser(bytes.NewReader(data)).Next()
require.NoError(t, err)
capsule, err := parseRouteAdvertisementCapsule(cr)
if tc.err != "" {
require.ErrorContains(t, err, tc.err)
return
}
require.NoError(t, err)
require.Equal(t, tc.routes, capsule.IPAddressRanges)
})
}
}
func TestAdvertiseRouteValidation(t *testing.T) {
tests := []struct {
name string
routes []IPRoute
err string
}{
{name: "invalid start IP", routes: []IPRoute{{EndIP: route4a.EndIP}}, err: "invalid IP address range"},
{name: "invalid end IP", routes: []IPRoute{{StartIP: route4a.StartIP}}, err: "invalid IP address range"},
{
name: "IPv6 zone",
routes: []IPRoute{{StartIP: netip.MustParseAddr("fe80::1%eth0"), EndIP: netip.MustParseAddr("fe80::2%eth0")}},
err: "invalid IP address range",
},
{name: "mixed IP versions", routes: []IPRoute{{StartIP: route4a.StartIP, EndIP: route6.EndIP}}, err: "mixes IP versions"},
{
name: "IPv4 and IPv4-mapped IPv6",
routes: []IPRoute{{StartIP: netip.MustParseAddr("10.0.0.1"), EndIP: netip.MustParseAddr("::ffff:10.0.0.2")}},
err: "mixes IP versions",
},
{name: "start after end", routes: []IPRoute{route4a, {StartIP: route4b.EndIP, EndIP: route4b.StartIP}}, err: "invalid route 1: start IP 10.0.0.20 is greater than end IP 10.0.0.10"},
}
tests = append(tests, routeOrderTests...)
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
t.Cleanup(func() { conn.Close() })
err := conn.AdvertiseRoute(tc.routes)
if tc.err != "" {
require.ErrorContains(t, err, tc.err)
conn.mu.Lock()
defer conn.mu.Unlock()
require.Empty(t, conn.queuedWrites)
require.Nil(t, conn.localRoutes)
return
}
require.NoError(t, err)
})
}
}
func TestReceiveMisorderedRouteAdvertisement(t *testing.T) {
toRead := make(chan []byte, 1)
conn := newProxiedConn(&mockStream{toRead: toRead})
t.Cleanup(func() { conn.Close() })
toRead <- (&routeAdvertisementCapsule{IPAddressRanges: []IPRoute{route6, route4a}}).append(nil)
ctx, cancel := context.WithTimeout(t.Context(), time.Second)
defer cancel()
_, err := conn.Routes(ctx)
require.ErrorIs(t, err, net.ErrClosed)
}
@@ -0,0 +1,95 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"crypto/rand"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"log"
"math/big"
"time"
"github.com/apernet/quic-go/http3"
)
var (
tlsConf *tls.Config
certPool *x509.CertPool
)
func generateCA() (*x509.Certificate, *rsa.PrivateKey, error) {
certTempl := &x509.Certificate{
SerialNumber: big.NewInt(2019),
Subject: pkix.Name{},
NotBefore: time.Now(),
NotAfter: time.Now().Add(24 * time.Hour),
IsCA: true,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth, x509.ExtKeyUsageServerAuth},
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign,
BasicConstraintsValid: true,
}
caPrivateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
return nil, nil, err
}
caBytes, err := x509.CreateCertificate(rand.Reader, certTempl, certTempl, &caPrivateKey.PublicKey, caPrivateKey)
if err != nil {
return nil, nil, err
}
ca, err := x509.ParseCertificate(caBytes)
if err != nil {
return nil, nil, err
}
return ca, caPrivateKey, nil
}
func generateLeafCert(ca *x509.Certificate, caPrivateKey *rsa.PrivateKey) (*x509.Certificate, *rsa.PrivateKey, error) {
certTempl := &x509.Certificate{
SerialNumber: big.NewInt(1),
DNSNames: []string{"localhost", "127.0.0.1"},
NotBefore: time.Now(),
NotAfter: time.Now().Add(24 * time.Hour),
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth, x509.ExtKeyUsageServerAuth},
KeyUsage: x509.KeyUsageDigitalSignature,
}
privKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
return nil, nil, err
}
certBytes, err := x509.CreateCertificate(rand.Reader, certTempl, ca, &privKey.PublicKey, caPrivateKey)
if err != nil {
return nil, nil, err
}
cert, err := x509.ParseCertificate(certBytes)
if err != nil {
return nil, nil, err
}
return cert, privKey, nil
}
func init() {
ca, caPrivateKey, err := generateCA()
if err != nil {
log.Fatal("failed to generate CA certificate:", err)
}
leafCert, leafPrivateKey, err := generateLeafCert(ca, caPrivateKey)
if err != nil {
log.Fatal("failed to generate leaf certificate:", err)
}
certPool = x509.NewCertPool()
certPool.AddCert(ca)
tlsConf = &tls.Config{
Certificates: []tls.Certificate{{
Certificate: [][]byte{leafCert.Raw},
PrivateKey: leafPrivateKey,
}},
NextProtos: []string{http3.NextProtoH3},
}
}
@@ -0,0 +1,23 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import "encoding/binary"
func calculateIPv4Checksum(header []byte) uint16 {
var sum uint32
for i := 0; i < len(header); i += 2 {
if i == 10 {
continue
}
sum += uint32(binary.BigEndian.Uint16(header[i : i+2]))
}
for (sum >> 16) > 0 {
sum = (sum & 0xffff) + (sum >> 16)
}
return ^uint16(sum)
}
@@ -0,0 +1,27 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"testing"
"github.com/stretchr/testify/require"
)
func TestIPv4ChecksumTestVector(t *testing.T) {
data := []byte{0x45, 0x00, 0x00, 0x73, 0x00, 0x00, 0x40, 0x00, 0x40, 0x11, 0xb8, 0x61, 0xc0, 0xa8, 0x00, 0x01, 0xc0, 0xa8, 0x00, 0xc7}
checksum := calculateIPv4Checksum(data)
require.Equal(t, uint16(0xb861), checksum)
}
func TestIPv4ChecksumWithOptions(t *testing.T) {
data := []byte{0x46, 0x00, 0x00, 0x77, 0x00, 0x00, 0x40, 0x00, 0x40, 0x11, 0x00, 0x00, 0xc0, 0xa8, 0x00, 0x01, 0xc0, 0xa8, 0x00, 0xc7, 0x94, 0x04, 0x00, 0x00}
checksum := calculateIPv4Checksum(data)
data[10], data[11] = byte(checksum>>8), byte(checksum)
require.True(t, ipv4ChecksumValid(data))
require.NotEqual(t, checksum, calculateIPv4Checksum(data[:20]), "the options must be covered")
}
@@ -0,0 +1,75 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"context"
"errors"
"fmt"
"net/http"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
)
type ClientConn struct {
clientConn *http3.ClientConn
}
func NewClientConn(conn *http3.ClientConn) *ClientConn {
return &ClientConn{clientConn: conn}
}
func (c *ClientConn) Dial(req *Request) (*Conn, *http.Response, error) {
httpReq := req.httpRequest()
if httpReq.URL == nil {
return nil, nil, errors.New("connect-ip: request URL is nil")
}
if httpReq.Host == "" && httpReq.URL.Host == "" {
return nil, nil, errors.New("connect-ip: request needs a host")
}
select {
case <-httpReq.Context().Done():
return nil, nil, context.Cause(httpReq.Context())
case <-c.clientConn.Context().Done():
return nil, nil, context.Cause(c.clientConn.Context())
case <-c.clientConn.ReceivedSettings():
}
settings := c.clientConn.Settings()
if !settings.EnableExtendedConnect {
return nil, nil, errors.New("connect-ip: server didn't enable Extended CONNECT")
}
if !settings.EnableDatagrams {
return nil, nil, errors.New("connect-ip: server didn't enable datagrams")
}
rstr, err := c.clientConn.OpenRequestStream(httpReq.Context())
if err != nil {
return nil, nil, fmt.Errorf("connect-ip: failed to open request stream: %w", err)
}
var keepStream bool
defer func() {
if !keepStream {
rstr.CancelRead(quic.StreamErrorCode(http3.ErrCodeNoError))
rstr.CancelWrite(quic.StreamErrorCode(http3.ErrCodeNoError))
}
}()
if err := rstr.SendRequestHeader(httpReq); err != nil {
return nil, nil, fmt.Errorf("connect-ip: failed to send request: %w", err)
}
rsp, err := rstr.ReadResponse()
if err != nil {
return nil, nil, fmt.Errorf("connect-ip: failed to read response: %w", err)
}
if rsp.StatusCode < 200 || rsp.StatusCode > 299 {
return nil, rsp, fmt.Errorf("connect-ip: server responded with %d", rsp.StatusCode)
}
keepStream = true
return newProxiedConn(rstr), rsp, nil
}
@@ -0,0 +1,101 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"context"
"net"
"net/http"
"testing"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestClientWaitForSettings(t *testing.T) {
conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0})
require.NoError(t, err)
ln, err := quic.Listen(conn, tlsConf, &quic.Config{EnableDatagrams: true})
require.NoError(t, err)
defer ln.Close()
h3conn := dialHTTP3(t, conn.LocalAddr().String())
ctx, cancel := context.WithTimeout(t.Context(), 100*time.Millisecond)
defer cancel()
req, err := NewRequest(ctx, "https://example.org/.well-known/masque/ip/")
require.NoError(t, err)
_, _, err = NewClientConn(h3conn).Dial(req)
require.ErrorIs(t, err, context.DeadlineExceeded)
}
func TestClientDatagramCheck(t *testing.T) {
s := http3.Server{
TLSConfig: tlsConf,
QUICConfig: &quic.Config{EnableDatagrams: true},
EnableDatagrams: false,
}
ln, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0})
require.NoError(t, err)
go func() { s.Serve(ln) }()
defer s.Close()
h3conn := dialHTTP3(t, ln.LocalAddr().String())
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
defer cancel()
req, err := NewRequest(ctx, "https://example.org/.well-known/masque/ip/")
require.NoError(t, err)
_, _, err = NewClientConn(h3conn).Dial(req)
require.ErrorContains(t, err, "connect-ip: server didn't enable datagrams")
}
func TestNewClientConnSharesHTTP3Connection(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
ln, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
require.NoError(t, err)
defer ln.Close()
url := "https://" + ln.LocalAddr().String()
mux := http.NewServeMux()
mux.HandleFunc("/connect-ip", func(w http.ResponseWriter, r *http.Request) {
req, err := ParseProxyRequest(r)
if !assert.NoError(t, err) {
w.WriteHeader(http.StatusBadRequest)
return
}
_, err = (&Proxy{}).Proxy(w, req)
assert.NoError(t, err)
})
mux.HandleFunc("GET /hello", func(http.ResponseWriter, *http.Request) {})
s := http3.Server{Handler: mux, TLSConfig: tlsConf, EnableDatagrams: true}
go func() { s.Serve(ln) }()
defer s.Close()
h3conn := dialHTTP3(t, ln.LocalAddr().String())
httpClient := &http.Client{Transport: h3conn, Timeout: time.Second}
checkHTTP := func() {
t.Helper()
rsp, err := httpClient.Get(url + "/hello")
require.NoError(t, err)
rsp.Body.Close()
require.Equal(t, http.StatusOK, rsp.StatusCode)
}
checkHTTP()
req, err := NewRequest(ctx, url+"/connect-ip")
require.NoError(t, err)
tunnel, rsp, err := NewClientConn(h3conn).Dial(req)
require.NoError(t, err)
require.Equal(t, http.StatusOK, rsp.StatusCode)
checkHTTP()
require.NoError(t, tunnel.Close())
checkHTTP()
}
+620
View File
@@ -0,0 +1,620 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"context"
"encoding/binary"
goerrors "errors"
"fmt"
"io"
"net"
"net/netip"
"slices"
"sync"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"github.com/apernet/quic-go/quicvarint"
"github.com/xtls/xray-core/common/errors"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
)
type CloseError struct {
Remote bool
}
func (e *CloseError) Error() string { return net.ErrClosed.Error() }
func (e *CloseError) Is(target error) bool { return target == net.ErrClosed }
const (
ipProtoICMP = 1
ipProtoICMPv6 = 58
)
type http3Stream interface {
io.ReadWriteCloser
StreamID() quic.StreamID
ReceiveDatagram(context.Context) ([]byte, error)
SendDatagram([]byte) error
CancelRead(quic.StreamErrorCode)
CancelWrite(quic.StreamErrorCode)
SetWriteDeadline(time.Time) error
}
var (
_ http3Stream = &http3.Stream{}
_ http3Stream = &http3.RequestStream{}
)
const maxQueuedCapsules = 128
var errCapsuleLimit = goerrors.New("connect-ip: capsule limit exceeded")
type streamWrite struct {
Data []byte
Fin bool
}
type Conn struct {
str http3Stream
writeNotify chan struct{}
writeDone chan error
assignedAddressUpdates chan []AssignedAddress
addressRequests chan *addressRequestCapsule
availableRouteUpdates chan []IPRoute
mu sync.Mutex
queuedWrites []streamWrite
peerAddresses []netip.Prefix
localRoutes []IPRoute
assignedAddresses []netip.Prefix
lastAddressRequestID AddressRequestID
closeChan chan struct{}
closeErr error
closeOnce sync.Once
closeResult error
datagramCapsuleOnce sync.Once
}
func newProxiedConn(str http3Stream) *Conn {
c := &Conn{
str: str,
writeNotify: make(chan struct{}, 1),
writeDone: make(chan error, 1),
assignedAddressUpdates: make(chan []AssignedAddress, maxQueuedCapsules),
addressRequests: make(chan *addressRequestCapsule, maxQueuedCapsules),
availableRouteUpdates: make(chan []IPRoute, 1),
closeChan: make(chan struct{}),
}
go func() {
err := c.readFromStream()
c.mu.Lock()
closing := c.closeErr != nil
if !closing {
c.closeErr = &CloseError{Remote: true}
close(c.closeChan)
if err != nil {
code := http3.ErrCodeMessageError
var streamErr *quic.StreamError
var h3Err *http3.Error
switch {
case goerrors.Is(err, errCapsuleLimit):
code = http3.ErrCodeExcessiveLoad
case goerrors.As(err, &streamErr) && streamErr.Remote, goerrors.As(err, &h3Err) && h3Err.Remote:
code = http3.ErrCodeRequestCanceled
}
c.str.CancelRead(quic.StreamErrorCode(code))
c.str.CancelWrite(quic.StreamErrorCode(code))
close(c.writeNotify)
} else {
c.queueFin()
}
}
c.mu.Unlock()
if err != nil && !closing {
errors.LogInfoInner(context.Background(), err, "reading capsules failed")
}
}()
go func() {
err := c.writeToStream()
if err != nil {
c.mu.Lock()
closing := c.closeErr != nil
if !closing {
c.closeErr = &CloseError{Remote: true}
close(c.closeChan)
c.str.CancelRead(quic.StreamErrorCode(http3.ErrCodeExcessiveLoad))
c.str.CancelWrite(quic.StreamErrorCode(http3.ErrCodeExcessiveLoad))
} else {
c.str.CancelWrite(quic.StreamErrorCode(http3.ErrCodeNoError))
}
c.mu.Unlock()
if !closing {
errors.LogInfoInner(context.Background(), err, "writing capsules failed")
}
}
c.writeDone <- err
close(c.writeDone)
}()
return c
}
func (c *Conn) AdvertiseRoute(routes []IPRoute) error {
for i, route := range routes {
err := route.validate()
if err == nil && i > 0 {
err = checkRouteOrder(routes[i-1], route)
}
if err != nil {
return fmt.Errorf("connect-ip: invalid route %d: %w", i, err)
}
}
c.mu.Lock()
if c.closeErr != nil {
err := c.closeErr
c.mu.Unlock()
return err
}
routes = slices.Clone(routes)
err := c.queueWrite(streamWrite{Data: (&routeAdvertisementCapsule{IPAddressRanges: routes}).append(nil)})
if err == nil {
c.localRoutes = routes
}
c.mu.Unlock()
if err != nil {
c.Close()
return err
}
return nil
}
func (c *Conn) RequestAddresses(prefixes []netip.Prefix) ([]AddressRequestID, error) {
if len(prefixes) == 0 {
return nil, goerrors.New("connect-ip: address request must contain at least one prefix")
}
for i, p := range prefixes {
if !p.IsValid() || p != p.Masked() {
return nil, fmt.Errorf("connect-ip: invalid requested prefix %d: %s", i, p)
}
}
c.mu.Lock()
if c.closeErr != nil {
err := c.closeErr
c.mu.Unlock()
return nil, err
}
ids := make([]AddressRequestID, len(prefixes))
for i := range ids {
ids[i] = c.lastAddressRequestID + AddressRequestID(i) + 1
}
capsule := &addressRequestCapsule{RequestIDs: ids, Prefixes: prefixes}
err := c.queueWrite(streamWrite{Data: capsule.append(nil)})
if err == nil {
c.lastAddressRequestID = ids[len(ids)-1]
}
c.mu.Unlock()
if err != nil {
c.Close()
return nil, err
}
return ids, nil
}
func (c *Conn) ReceiveAddressAssignment(ctx context.Context) ([]AssignedAddress, error) {
select {
case <-ctx.Done():
return nil, ctx.Err()
case assignment := <-c.assignedAddressUpdates:
return assignment, nil
case <-c.closeChan:
select {
case assignment := <-c.assignedAddressUpdates:
return assignment, nil
default:
return nil, c.closeErr
}
}
}
func (c *Conn) ReceiveAddressRequest(ctx context.Context) (*AddressRequest, error) {
var requested *addressRequestCapsule
select {
case <-ctx.Done():
return nil, ctx.Err()
case requested = <-c.addressRequests:
case <-c.closeChan:
select {
case requested = <-c.addressRequests:
default:
return nil, c.closeErr
}
}
return newAddressRequest(c, requested), nil
}
func (c *Conn) AssignAddresses(prefixes []netip.Prefix) error {
capsule := &addressAssignCapsule{}
if prefixes != nil {
capsule.AssignedAddresses = make([]AssignedAddress, len(prefixes))
for i, p := range prefixes {
capsule.AssignedAddresses[i] = AssignedAddress{IPPrefix: p}
}
}
return c.sendAddressAssignment(capsule, true)
}
func (c *Conn) sendAddressAssignment(capsule *addressAssignCapsule, restrictPeer bool) error {
c.mu.Lock()
if c.closeErr != nil {
err := c.closeErr
c.mu.Unlock()
return err
}
if err := c.queueWrite(streamWrite{Data: capsule.append(nil)}); err != nil {
c.mu.Unlock()
c.Close()
return err
}
if !restrictPeer && c.peerAddresses == nil {
c.mu.Unlock()
return nil
}
var prefixes []netip.Prefix
if capsule.AssignedAddresses != nil {
prefixes = make([]netip.Prefix, 0, len(capsule.AssignedAddresses))
}
for _, assigned := range capsule.AssignedAddresses {
if !assigned.Rejected() {
prefixes = append(prefixes, assigned.IPPrefix)
}
}
c.peerAddresses = prefixes
c.mu.Unlock()
return nil
}
func (c *Conn) queueWrite(w streamWrite) error {
if len(c.queuedWrites) >= maxQueuedCapsules {
c.closeErr = &CloseError{Remote: false}
close(c.closeChan)
c.str.CancelRead(quic.StreamErrorCode(http3.ErrCodeExcessiveLoad))
c.str.CancelWrite(quic.StreamErrorCode(http3.ErrCodeExcessiveLoad))
close(c.writeNotify)
return goerrors.New("connect-ip: capsule queue full")
}
c.queuedWrites = append(c.queuedWrites, w)
c.notifyWriter()
return nil
}
func (c *Conn) queueFin() {
c.str.SetWriteDeadline(time.Now())
c.queuedWrites = append(c.queuedWrites, streamWrite{Fin: true})
c.notifyWriter()
}
func (c *Conn) notifyWriter() {
select {
case c.writeNotify <- struct{}{}:
default:
}
}
func queueLatest[T any](ch chan T, value T) {
for {
select {
case ch <- value:
return
case <-ch:
}
}
}
func (c *Conn) Routes(ctx context.Context) ([]IPRoute, error) {
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-c.closeChan:
return nil, c.closeErr
case routes := <-c.availableRouteUpdates:
return routes, nil
}
}
func (c *Conn) readFromStream() error {
p := http3.NewCapsuleParser(c.str)
for {
t, cr, err := p.Next()
if goerrors.Is(err, io.EOF) {
return nil
}
if err != nil {
return err
}
switch t {
case capsuleTypeAddressAssign:
capsule, err := parseAddressAssignCapsule(cr)
if err != nil {
return err
}
prefixes := make([]netip.Prefix, 0, len(capsule.AssignedAddresses))
for _, assigned := range capsule.AssignedAddresses {
if !assigned.Rejected() {
prefixes = append(prefixes, assigned.IPPrefix)
}
}
c.mu.Lock()
c.assignedAddresses = prefixes
c.mu.Unlock()
select {
case c.assignedAddressUpdates <- capsule.AssignedAddresses:
default:
return fmt.Errorf("%w: address assignment queue full", errCapsuleLimit)
}
case capsuleTypeAddressRequest:
capsule, err := parseAddressRequestCapsule(cr)
if err != nil {
return err
}
select {
case c.addressRequests <- capsule:
default:
return fmt.Errorf("%w: address request queue full", errCapsuleLimit)
}
case capsuleTypeRouteAdvertisement:
capsule, err := parseRouteAdvertisementCapsule(cr)
if err != nil {
return err
}
queueLatest(c.availableRouteUpdates, capsule.IPAddressRanges)
case capsuleTypeDatagram:
c.datagramCapsuleOnce.Do(func() {
errors.LogWarning(context.Background(), "connect-ip: dropping IP packets sent in DATAGRAM capsules, only QUIC DATAGRAM frames are supported")
})
if err := cr.Discard(); err != nil {
return err
}
default:
if err := cr.Discard(); err != nil {
return err
}
}
}
}
func (c *Conn) writeToStream() error {
for range c.writeNotify {
for {
c.mu.Lock()
if len(c.queuedWrites) == 0 {
c.mu.Unlock()
break
}
w := c.queuedWrites[0]
c.queuedWrites[0] = streamWrite{}
c.queuedWrites = c.queuedWrites[1:]
c.mu.Unlock()
if w.Fin {
return c.str.Close()
}
if _, err := c.str.Write(w.Data); err != nil {
return err
}
}
}
return c.closeErr
}
func (c *Conn) ReadPacket(b []byte) (int, error) {
for {
select {
case <-c.closeChan:
return 0, c.closeErr
default:
}
data, err := c.str.ReceiveDatagram(context.Background())
if err != nil {
select {
case <-c.closeChan:
return 0, c.closeErr
default:
return 0, err
}
}
contextID, n, err := quicvarint.Parse(data)
if err != nil {
errors.LogDebugInner(context.Background(), err, "dropping malformed datagram")
continue
}
if contextID != 0 {
continue
}
packet := data[n:]
if err := c.handleIncomingProxiedPacket(packet); err != nil {
errors.LogDebugInner(context.Background(), err, "dropping proxied packet")
continue
}
if len(packet) > len(b) {
return 0, io.ErrShortBuffer
}
return copy(b, packet), nil
}
}
func (c *Conn) handleIncomingProxiedPacket(data []byte) error {
if len(data) == 0 {
return goerrors.New("connect-ip: empty packet")
}
var src, dst netip.Addr
var ipProto uint8
switch v := ipVersion(data); v {
default:
return fmt.Errorf("connect-ip: unknown IP versions: %d", v)
case 4:
if len(data) < ipv4.HeaderLen {
return fmt.Errorf("connect-ip: malformed datagram: too short")
}
src = netip.AddrFrom4([4]byte(data[12:16]))
dst = netip.AddrFrom4([4]byte(data[16:20]))
ipProto = data[9]
case 6:
if len(data) < ipv6.HeaderLen {
return fmt.Errorf("connect-ip: malformed datagram: too short")
}
src = netip.AddrFrom16([16]byte(data[8:24]))
dst = netip.AddrFrom16([16]byte(data[24:40]))
ipProto = data[6]
}
c.mu.Lock()
assignedAddresses := c.assignedAddresses
localRoutes := c.localRoutes
peerAddresses := c.peerAddresses
c.mu.Unlock()
if peerAddresses != nil {
if !slices.ContainsFunc(peerAddresses, func(p netip.Prefix) bool { return p.Contains(src) }) {
return fmt.Errorf("connect-ip: datagram source address not allowed: %s", src)
}
}
var isAllowedDst bool
if len(assignedAddresses) > 0 {
isAllowedDst = slices.ContainsFunc(assignedAddresses, func(p netip.Prefix) bool { return p.Contains(dst) })
}
if !isAllowedDst {
isAllowedDst = slices.ContainsFunc(localRoutes, func(r IPRoute) bool {
if r.StartIP.Compare(dst) > 0 || dst.Compare(r.EndIP) > 0 {
return false
}
if (ipVersion(data) == 4 && ipProto == ipProtoICMP) || (ipVersion(data) == 6 && ipProto == ipProtoICMPv6) {
return true
}
return r.IPProtocol == 0 || r.IPProtocol == ipProto
})
}
if !isAllowedDst {
return fmt.Errorf("connect-ip: datagram destination address / protocol not allowed: %s (protocol: %d)", dst, ipProto)
}
return nil
}
func (c *Conn) WritePacket(b []byte) (icmp []byte, err error) {
select {
case <-c.closeChan:
return nil, c.closeErr
default:
}
data, err := c.composeDatagram(b)
if err != nil {
errors.LogDebugInner(context.Background(), err, "dropping proxied packet (", len(b), " bytes) that can't be proxied")
return nil, nil
}
if err := c.str.SendDatagram(data); err != nil {
if tooLarge, ok := goerrors.AsType[*quic.DatagramTooLargeError](err); ok {
icmpPacket, err := composeICMPTooLargePacket(b, int(tooLarge.MaxDatagramPayloadSize)-c.datagramOverhead())
if err != nil {
if goerrors.Is(err, ErrMTUTooSmall) {
return nil, err
}
errors.LogDebugInner(context.Background(), err, "failed to compose ICMP Packet Too Big")
}
return icmpPacket, nil
}
select {
case <-c.closeChan:
return nil, c.closeErr
default:
return nil, err
}
}
return nil, nil
}
func (c *Conn) composeDatagram(b []byte) ([]byte, error) {
if len(b) == 0 {
return nil, goerrors.New("connect-ip: empty packet")
}
switch v := ipVersion(b); v {
default:
return nil, fmt.Errorf("connect-ip: unknown IP versions: %d", v)
case 4:
if len(b) < ipv4.HeaderLen {
return nil, fmt.Errorf("connect-ip: IPv4 packet too short")
}
hdrLen := int(b[0]&0x0f) << 2
totalLen := int(binary.BigEndian.Uint16(b[2:4]))
if hdrLen < ipv4.HeaderLen || hdrLen > totalLen || totalLen > len(b) {
return nil, fmt.Errorf("connect-ip: malformed IPv4 header: header length %d, total length %d, packet length %d", hdrLen, totalLen, len(b))
}
ttl := b[8]
if ttl <= 1 {
return nil, fmt.Errorf("connect-ip: datagram TTL too small: %d", ttl)
}
b[8]--
binary.BigEndian.PutUint16(b[10:12], calculateIPv4Checksum(b[:hdrLen]))
case 6:
if len(b) < ipv6.HeaderLen {
return nil, fmt.Errorf("connect-ip: IPv6 packet too short")
}
hopLimit := b[7]
if hopLimit <= 1 {
return nil, fmt.Errorf("connect-ip: datagram Hop Limit too small: %d", hopLimit)
}
b[7]--
}
data := make([]byte, 0, len(contextIDZero)+len(b))
data = append(data, contextIDZero...)
data = append(data, b...)
return data, nil
}
func (c *Conn) datagramOverhead() int {
return quicvarint.Len(uint64(c.str.StreamID()/4)) + len(contextIDZero)
}
func (c *Conn) MaxPacketSize() int {
select {
case <-c.closeChan:
return 0
default:
}
err := c.str.SendDatagram(make([]byte, 1<<16))
tooLarge, ok := goerrors.AsType[*quic.DatagramTooLargeError](err)
if !ok {
return 0
}
return max(0, int(tooLarge.MaxDatagramPayloadSize)-c.datagramOverhead())
}
func (c *Conn) Close() error {
c.closeOnce.Do(func() {
c.mu.Lock()
if c.closeErr == nil {
c.closeErr = &CloseError{Remote: false}
close(c.closeChan)
c.queueFin()
}
c.mu.Unlock()
c.closeResult = <-c.writeDone
c.str.CancelRead(quic.StreamErrorCode(http3.ErrCodeNoError))
})
return c.closeResult
}
func ipVersion(b []byte) uint8 { return b[0] >> 4 }
@@ -0,0 +1,706 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"bytes"
"context"
"encoding/binary"
"errors"
"io"
"net"
"net/netip"
"sync"
"testing"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"github.com/apernet/quic-go/quicvarint"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/net/icmp"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
)
var ipv6Header = []byte{
0x60, 0x00, 0x00, 0x00,
0x00, 0x20, 59, 64,
0x20, 0x01, 0x0d, 0xb8, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01,
0x20, 0x01, 0x0d, 0xb8, 0x85, 0xa3, 0x08, 0xd3, 0x13, 0x19, 0x8a, 0x2e, 0x03, 0x70, 0x73, 0x48,
}
var (
testSrc4 = netip.MustParseAddr("192.0.2.1")
testDst4 = netip.MustParseAddr("198.51.100.1")
testSrc6 = netip.MustParseAddr("2001:db8::1")
testDst6 = netip.MustParseAddr("2001:db8:1::1")
)
func ipv4Packet(ttl, proto uint8, src, dst netip.Addr, options, payload []byte) []byte {
hdrLen := ipv4.HeaderLen + len(options)
b := make([]byte, hdrLen, hdrLen+len(payload))
b[0] = 4<<4 | byte(hdrLen>>2)
binary.BigEndian.PutUint16(b[2:4], uint16(hdrLen+len(payload)))
b[8] = ttl
b[9] = proto
copy(b[12:16], src.AsSlice())
copy(b[16:20], dst.AsSlice())
copy(b[ipv4.HeaderLen:], options)
return append(b, payload...)
}
func ipv6Packet(hopLimit, nextHeader uint8, src, dst netip.Addr, payload []byte) []byte {
b := make([]byte, ipv6.HeaderLen, ipv6.HeaderLen+len(payload))
b[0] = 6 << 4
binary.BigEndian.PutUint16(b[4:6], uint16(len(payload)))
b[6] = nextHeader
b[7] = hopLimit
copy(b[8:24], src.AsSlice())
copy(b[24:40], dst.AsSlice())
return append(b, payload...)
}
func ipv4ChecksumValid(header []byte) bool {
var sum uint32
for i := 0; i+1 < len(header); i += 2 {
sum += uint32(binary.BigEndian.Uint16(header[i:]))
}
for sum > 0xffff {
sum = sum&0xffff + sum>>16
}
return sum == 0xffff
}
type mockStream struct {
streamID quic.StreamID
reading []byte
toRead <-chan []byte
datagrams <-chan []byte
maxDatagramPayloadSize int
sendDatagramErr error
sent [][]byte
writeStarted chan struct{}
written chan<- []byte
readErr error
mu sync.Mutex
cancelWriteCodes []quic.StreamErrorCode
}
func (m *mockStream) cancelWriteCode() (quic.StreamErrorCode, bool) {
m.mu.Lock()
defer m.mu.Unlock()
if len(m.cancelWriteCodes) == 0 {
return 0, false
}
return m.cancelWriteCodes[0], true
}
var _ http3Stream = &mockStream{}
func (m *mockStream) StreamID() quic.StreamID { return m.streamID }
func (m *mockStream) Read(p []byte) (int, error) {
if len(m.reading) == 0 && m.readErr != nil {
return 0, m.readErr
}
if len(m.reading) == 0 {
m.reading = <-m.toRead
}
n := copy(p, m.reading)
m.reading = m.reading[n:]
return n, nil
}
func (m *mockStream) CancelRead(quic.StreamErrorCode) {}
func (m *mockStream) Write(p []byte) (int, error) {
if m.writeStarted != nil {
close(m.writeStarted)
m.writeStarted = nil
}
if m.written != nil {
m.written <- bytes.Clone(p)
}
return len(p), nil
}
func (m *mockStream) Close() error { return nil }
func (m *mockStream) CancelWrite(code quic.StreamErrorCode) {
m.mu.Lock()
defer m.mu.Unlock()
m.cancelWriteCodes = append(m.cancelWriteCodes, code)
}
func (m *mockStream) SetWriteDeadline(time.Time) error { return nil }
func (m *mockStream) SendDatagram(data []byte) error {
if m.sendDatagramErr != nil {
return m.sendDatagramErr
}
if size := quicvarint.Len(uint64(m.streamID/4)) + len(data); m.maxDatagramPayloadSize > 0 && size > m.maxDatagramPayloadSize {
return &quic.DatagramTooLargeError{MaxDatagramPayloadSize: int64(m.maxDatagramPayloadSize)}
}
m.sent = append(m.sent, bytes.Clone(data))
return nil
}
func (m *mockStream) ReceiveDatagram(ctx context.Context) ([]byte, error) {
select {
case <-ctx.Done():
return nil, ctx.Err()
case data, ok := <-m.datagrams:
if !ok {
return nil, io.EOF
}
return data, nil
}
}
func TestCapsuleWriteQueueLimit(t *testing.T) {
writes := make(chan []byte)
writeStarted := make(chan struct{})
conn := newProxiedConn(&mockStream{
writeStarted: writeStarted,
written: writes,
})
t.Cleanup(func() { conn.Close() })
require.NoError(t, conn.AssignAddresses(nil))
select {
case <-writeStarted:
case <-time.After(time.Second):
t.Fatal("capsule write did not start")
}
for range maxQueuedCapsules {
require.NoError(t, conn.AssignAddresses(nil))
}
go func() {
conn.Routes(context.Background())
for range maxQueuedCapsules + 1 {
<-writes
}
}()
require.ErrorContains(t, conn.AssignAddresses(nil), "capsule queue full")
require.ErrorIs(t, conn.AssignAddresses(nil), net.ErrClosed)
}
func TestCapsuleReceiveQueueLimit(t *testing.T) {
for _, name := range []string{"assignments", "requests"} {
t.Run(name, func(t *testing.T) {
var data []byte
for i := range maxQueuedCapsules + 1 {
if name == "assignments" {
data = (&addressAssignCapsule{}).append(data)
} else {
data = (&addressRequestCapsule{
RequestIDs: []AddressRequestID{AddressRequestID(i + 1)},
Prefixes: []netip.Prefix{netip.MustParsePrefix("192.0.2.1/32")},
}).append(data)
}
}
conn := newProxiedConn(&mockStream{reading: data})
t.Cleanup(func() { conn.Close() })
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_, err := conn.Routes(ctx)
require.ErrorIs(t, err, net.ErrClosed)
})
}
}
func TestAbortErrorCode(t *testing.T) {
var overflow []byte
for range maxQueuedCapsules + 1 {
overflow = (&addressAssignCapsule{}).append(overflow)
}
misordered := (&routeAdvertisementCapsule{IPAddressRanges: []IPRoute{
{StartIP: netip.MustParseAddr("192.0.2.2"), EndIP: netip.MustParseAddr("192.0.2.1")},
}}).append(nil)
for _, c := range []struct {
name string
str *mockStream
code http3.ErrCode
}{
{"malformed capsule", &mockStream{reading: misordered}, http3.ErrCodeMessageError},
{"queue limit", &mockStream{reading: overflow}, http3.ErrCodeExcessiveLoad},
{"reset by peer", &mockStream{readErr: &quic.StreamError{ErrorCode: quic.StreamErrorCode(http3.ErrCodeNoError), Remote: true}}, http3.ErrCodeRequestCanceled},
{"reset by peer on a request stream", &mockStream{readErr: &http3.Error{ErrorCode: http3.ErrCodeNoError, Remote: true}}, http3.ErrCodeRequestCanceled},
} {
t.Run(c.name, func(t *testing.T) {
conn := newProxiedConn(c.str)
t.Cleanup(func() { conn.Close() })
require.Eventually(t, func() bool {
_, ok := c.str.cancelWriteCode()
return ok
}, time.Second, time.Millisecond)
code, _ := c.str.cancelWriteCode()
require.Equal(t, quic.StreamErrorCode(c.code), code)
})
}
}
func TestIncomingDatagrams(t *testing.T) {
t.Run("empty packets", func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
require.ErrorContains(t,
conn.handleIncomingProxiedPacket([]byte{}),
"connect-ip: empty packet",
)
})
t.Run("invalid IP version", func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
data := make([]byte, 20)
data[0] = 5 << 4
require.ErrorContains(t,
conn.handleIncomingProxiedPacket(data),
"connect-ip: unknown IP versions: 5",
)
})
t.Run("IPv4 packet too short", func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
data, err := (&ipv4.Header{
Src: net.IPv4(1, 2, 3, 4),
Dst: net.IPv4(159, 70, 42, 98),
Len: 20,
Checksum: 89,
}).Marshal()
require.NoError(t, err)
require.ErrorContains(t,
conn.handleIncomingProxiedPacket(data[:ipv4.HeaderLen-1]),
"connect-ip: malformed datagram: too short",
)
})
t.Run("IPv6 packet too short", func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
require.ErrorContains(t,
conn.handleIncomingProxiedPacket(ipv6Header[:ipv6.HeaderLen-1]),
"connect-ip: malformed datagram: too short",
)
})
t.Run("invalid source address", func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
require.NoError(t, conn.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("192.168.0.10/32")}))
hdr := &ipv4.Header{
Src: net.IPv4(192, 168, 0, 11),
Dst: net.IPv4(159, 70, 42, 98),
Len: 20,
Checksum: 89,
}
data, err := hdr.Marshal()
require.NoError(t, err)
require.ErrorContains(t,
conn.handleIncomingProxiedPacket(data),
"connect-ip: datagram source address not allowed: 192.168.0.11",
)
})
t.Run("invalid destination address", func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
require.NoError(t, conn.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("192.168.0.10/32")}))
require.NoError(t, conn.AdvertiseRoute([]IPRoute{
{StartIP: netip.MustParseAddr("10.0.0.0"), EndIP: netip.MustParseAddr("10.1.2.3")},
}))
hdr := &ipv4.Header{
Src: net.IPv4(192, 168, 0, 10),
Dst: net.IPv4(10, 1, 2, 3),
Len: 20,
Checksum: 89,
}
data, err := hdr.Marshal()
require.NoError(t, err)
require.NoError(t, conn.handleIncomingProxiedPacket(data))
hdr.Dst = net.IPv4(10, 1, 2, 4)
data, err = hdr.Marshal()
require.NoError(t, err)
require.ErrorContains(t,
conn.handleIncomingProxiedPacket(data),
"connect-ip: datagram destination address / protocol not allowed: 10.1.2.4 (protocol: 0)",
)
})
t.Run("invalid IP protocol", func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
require.NoError(t, conn.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("192.168.0.10/32")}))
require.NoError(t, conn.AdvertiseRoute([]IPRoute{
{StartIP: netip.MustParseAddr("10.0.0.0"), EndIP: netip.MustParseAddr("10.1.2.3"), IPProtocol: 42},
}))
hdr := &ipv4.Header{
Src: net.IPv4(192, 168, 0, 10),
Dst: net.IPv4(10, 1, 2, 3),
Len: 20,
Checksum: 89,
Protocol: 42,
}
data, err := hdr.Marshal()
require.NoError(t, err)
require.NoError(t, conn.handleIncomingProxiedPacket(data))
hdr.Protocol = 41
data, err = hdr.Marshal()
require.NoError(t, err)
require.ErrorContains(t,
conn.handleIncomingProxiedPacket(data),
"connect-ip: datagram destination address / protocol not allowed: 10.1.2.3 (protocol: 41)",
)
hdr.Protocol = ipProtoICMP
data, err = hdr.Marshal()
require.NoError(t, err)
require.NoError(t, conn.handleIncomingProxiedPacket(data))
})
t.Run("packet from assigned address", func(t *testing.T) {
readChan := make(chan []byte, 1)
conn := newProxiedConn(&mockStream{toRead: readChan})
hdr := &ipv4.Header{
Src: net.IPv4(159, 70, 42, 98),
Dst: net.IPv4(192, 168, 0, 10),
Len: 20,
Checksum: 89,
}
data, err := hdr.Marshal()
require.NoError(t, err)
require.Error(t, conn.handleIncomingProxiedPacket(data), "connect-ip: datagram destination address")
readChan <- (&addressAssignCapsule{
AssignedAddresses: []AssignedAddress{{IPPrefix: netip.MustParsePrefix("192.168.0.10/32")}},
}).append(nil)
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_, err = conn.ReceiveAddressAssignment(ctx)
require.NoError(t, err)
require.NoError(t, conn.handleIncomingProxiedPacket(data))
})
}
func TestSkipUnknownCapsule(t *testing.T) {
for _, typ := range []http3.CapsuleType{42, capsuleTypeDatagram} {
readChan := make(chan []byte, 1)
conn := newProxiedConn(&mockStream{toRead: readChan})
data := quicvarint.Append(nil, uint64(typ))
data = quicvarint.Append(data, 3)
data = append(data, "foo"...)
data = (&addressAssignCapsule{
AssignedAddresses: []AssignedAddress{{IPPrefix: netip.MustParsePrefix("192.168.0.10/32")}},
}).append(data)
readChan <- data
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
assigned, err := conn.ReceiveAddressAssignment(ctx)
cancel()
require.NoError(t, err, "capsule type %d", typ)
require.Equal(t, []AssignedAddress{{IPPrefix: netip.MustParsePrefix("192.168.0.10/32")}}, assigned)
conn.Close()
}
}
func FuzzIncomingDatagram(f *testing.F) {
conn := newProxiedConn(&mockStream{})
require.NoError(f, conn.AssignAddresses([]netip.Prefix{
netip.MustParsePrefix("192.168.0.0/16"),
netip.MustParsePrefix("2001:db8::0/64"),
}))
require.NoError(f, conn.AdvertiseRoute([]IPRoute{
{StartIP: netip.MustParseAddr("10.0.0.0"), EndIP: netip.MustParseAddr("10.1.2.3"), IPProtocol: 42},
{StartIP: netip.MustParseAddr("2001:db8:1::"), EndIP: netip.MustParseAddr("2001:db8:1::ffff"), IPProtocol: 42},
}))
ipv4Header, err := (&ipv4.Header{
Src: net.IPv4(1, 2, 3, 4),
Dst: net.IPv4(159, 70, 42, 98),
Len: 20,
Checksum: 89,
}).Marshal()
require.NoError(f, err)
f.Add(ipv4Header)
f.Add(ipv6Header)
f.Fuzz(func(t *testing.T, data []byte) {
conn.handleIncomingProxiedPacket(data)
})
}
func TestSendingDatagrams(t *testing.T) {
t.Run("invalid IP version", func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
data := make([]byte, 20)
data[0] = 5 << 4
_, err := conn.composeDatagram(data)
require.ErrorContains(t, err, "connect-ip: unknown IP versions: 5")
})
t.Run("IPv4 packet too short", func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
data, err := (&ipv4.Header{
Src: net.IPv4(1, 2, 3, 4),
Dst: net.IPv4(159, 70, 42, 98),
Len: 20,
Checksum: 89,
}).Marshal()
require.NoError(t, err)
_, err = conn.composeDatagram(data[:ipv4.HeaderLen-1])
require.ErrorContains(t, err, "connect-ip: IPv4 packet too short")
})
t.Run("IPv6 packet too short", func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
_, err := conn.composeDatagram(ipv6Header[:ipv6.HeaderLen-1])
require.ErrorContains(t, err, "connect-ip: IPv6 packet too short")
})
}
func TestWritePacketDropsWithoutSending(t *testing.T) {
setIHL := func(b []byte, ihl byte) []byte { b[0] = 4<<4 | ihl; return b }
setTotalLen := func(b []byte, l uint16) []byte { binary.BigEndian.PutUint16(b[2:4], l); return b }
payload := make([]byte, 20)
for _, tc := range []struct {
name string
packet []byte
}{
{"nil", nil},
{"empty", []byte{}},
{"IPv4 TTL 1", ipv4Packet(1, 17, testSrc4, testDst4, nil, payload)},
{"IPv4 TTL 0", ipv4Packet(0, 17, testSrc4, testDst4, nil, payload)},
{"IPv6 Hop Limit 1", ipv6Packet(1, 17, testSrc6, testDst6, payload)},
{"IPv6 Hop Limit 0", ipv6Packet(0, 17, testSrc6, testDst6, payload)},
{"IPv4 IHL below 5", setIHL(ipv4Packet(64, 17, testSrc4, testDst4, nil, payload), 4)},
{"IPv4 IHL beyond total length", setIHL(ipv4Packet(64, 17, testSrc4, testDst4, nil, payload[:8]), 8)},
{"IPv4 IHL beyond packet", setTotalLen(setIHL(ipv4Packet(64, 17, testSrc4, testDst4, nil, payload), 15), 60)},
{"IPv4 total length beyond packet", setTotalLen(ipv4Packet(64, 17, testSrc4, testDst4, nil, payload), 41)},
{"IPv4 total length below header", setTotalLen(ipv4Packet(64, 17, testSrc4, testDst4, nil, payload), 19)},
} {
t.Run(tc.name, func(t *testing.T) {
str := &mockStream{}
conn := newProxiedConn(str)
t.Cleanup(func() { conn.Close() })
orig := bytes.Clone(tc.packet)
icmpPacket, err := conn.WritePacket(tc.packet)
require.NoError(t, err)
require.Nil(t, icmpPacket)
require.Empty(t, str.sent)
require.Equal(t, orig, tc.packet, "dropped packets must not be modified")
})
}
}
func TestWritePacketIPv4Checksum(t *testing.T) {
for _, tc := range []struct {
name string
options []byte
}{
{"no options", nil},
{"Router Alert option", []byte{0x94, 0x04, 0x00, 0x00}},
{"maximum header length", bytes.Repeat([]byte{0x01}, 40)},
} {
t.Run(tc.name, func(t *testing.T) {
str := &mockStream{}
conn := newProxiedConn(str)
t.Cleanup(func() { conn.Close() })
packet := ipv4Packet(64, 17, testSrc4, testDst4, tc.options, []byte("foobar"))
icmpPacket, err := conn.WritePacket(packet)
require.NoError(t, err)
require.Nil(t, icmpPacket)
require.Len(t, str.sent, 1)
require.Equal(t, contextIDZero, str.sent[0][:len(contextIDZero)])
sent := str.sent[0][len(contextIDZero):]
require.Len(t, sent, len(packet))
require.Equal(t, uint8(63), sent[8])
require.True(t, ipv4ChecksumValid(sent[:ipv4.HeaderLen+len(tc.options)]))
})
}
}
func TestWritePacketTooLarge(t *testing.T) {
for _, tc := range []struct {
name string
streamID quic.StreamID
maxPayloadSize int
ipv6 bool
wantMTU int
wantMTUTooSmall bool
}{
{name: "IPv4", maxPayloadSize: 1200, wantMTU: 1198},
{name: "IPv4, 2-byte Quarter Stream ID", streamID: 4 * 64, maxPayloadSize: 1200, wantMTU: 1197},
{name: "IPv4 minimum MTU", maxPayloadSize: 70, wantMTU: 68},
{name: "IPv4 below minimum MTU", maxPayloadSize: 69, wantMTU: 67, wantMTUTooSmall: true},
{name: "IPv6", maxPayloadSize: 1400, ipv6: true, wantMTU: 1398},
{name: "IPv6, 4-byte Quarter Stream ID", streamID: 4 * 20000, maxPayloadSize: 1400, ipv6: true, wantMTU: 1395},
{name: "IPv6 minimum MTU", maxPayloadSize: 1282, ipv6: true, wantMTU: 1280},
{name: "IPv6 below minimum MTU", maxPayloadSize: 1281, ipv6: true, wantMTU: 1279, wantMTUTooSmall: true},
} {
t.Run(tc.name, func(t *testing.T) {
str := &mockStream{streamID: tc.streamID, maxDatagramPayloadSize: tc.maxPayloadSize}
conn := newProxiedConn(str)
t.Cleanup(func() { conn.Close() })
packetOfSize := func(size int) []byte {
if tc.ipv6 {
return ipv6Packet(64, 17, testSrc6, testDst6, make([]byte, size-ipv6.HeaderLen))
}
return ipv4Packet(64, 17, testSrc4, testDst4, nil, make([]byte, size-ipv4.HeaderLen))
}
require.Equal(t, tc.wantMTU, conn.MaxPacketSize())
icmpPacket, err := conn.WritePacket(packetOfSize(tc.wantMTU))
require.NoError(t, err)
require.Nil(t, icmpPacket)
require.Len(t, str.sent, 1)
icmpPacket, err = conn.WritePacket(packetOfSize(tc.wantMTU + 1))
if tc.wantMTUTooSmall {
require.ErrorIs(t, err, ErrMTUTooSmall)
require.Nil(t, icmpPacket)
return
}
require.NoError(t, err)
if tc.ipv6 {
msg, err := icmp.ParseMessage(ipProtoICMPv6, icmpPacket[ipv6.HeaderLen:])
require.NoError(t, err)
require.Equal(t, ipv6.ICMPTypePacketTooBig, msg.Type)
require.Equal(t, tc.wantMTU, msg.Body.(*icmp.PacketTooBig).MTU)
} else {
msg := icmpPacket[ipv4.HeaderLen:]
require.Equal(t, []byte{3, 4}, msg[:2])
require.Equal(t, uint16(tc.wantMTU), binary.BigEndian.Uint16(msg[6:8]))
}
})
}
}
func TestMaxPacketSize(t *testing.T) {
t.Run("sends nothing", func(t *testing.T) {
str := &mockStream{streamID: 8, maxDatagramPayloadSize: 1350}
conn := newProxiedConn(str)
t.Cleanup(func() { conn.Close() })
require.Equal(t, 1348, conn.MaxPacketSize())
require.Empty(t, str.sent)
})
t.Run("datagrams unsupported", func(t *testing.T) {
conn := newProxiedConn(&mockStream{sendDatagramErr: errors.New("datagram support disabled")})
t.Cleanup(func() { conn.Close() })
require.Zero(t, conn.MaxPacketSize())
})
t.Run("closed", func(t *testing.T) {
conn := newProxiedConn(&mockStream{maxDatagramPayloadSize: 1350})
require.NoError(t, conn.Close())
require.Zero(t, conn.MaxPacketSize())
})
}
func TestReadPacketDropsMalformedDatagrams(t *testing.T) {
packet := ipv4Packet(64, 17, testSrc4, testDst4, nil, []byte("foobar"))
for _, tc := range []struct {
name string
datagram []byte
}{
{"empty", []byte{}},
{"truncated Context ID", []byte{0x40}},
{"unknown Context ID", append([]byte{0x02}, packet...)},
{"empty IP packet", []byte{0x00}},
{"invalid IP packet", []byte{0x00, 0x50}},
} {
t.Run(tc.name, func(t *testing.T) {
datagrams := make(chan []byte, 2)
conn := newProxiedConn(&mockStream{datagrams: datagrams})
t.Cleanup(func() { conn.Close() })
require.NoError(t, conn.AdvertiseRoute([]IPRoute{
{StartIP: netip.IPv4Unspecified(), EndIP: netip.MustParseAddr("255.255.255.255")},
}))
datagrams <- tc.datagram
datagrams <- append(bytes.Clone(contextIDZero), packet...)
b := make([]byte, 1500)
n, err := conn.ReadPacket(b)
require.NoError(t, err)
require.Equal(t, packet, b[:n])
})
}
}
func TestReadPacketShortBuffer(t *testing.T) {
datagrams := make(chan []byte, 3)
conn := newProxiedConn(&mockStream{datagrams: datagrams})
t.Cleanup(func() { conn.Close() })
require.NoError(t, conn.AdvertiseRoute([]IPRoute{
{StartIP: netip.IPv4Unspecified(), EndIP: netip.MustParseAddr("255.255.255.255")},
}))
packet := ipv4Packet(64, 17, testSrc4, testDst4, nil, []byte("foobar"))
for range 3 {
datagrams <- append(bytes.Clone(contextIDZero), packet...)
}
n, err := conn.ReadPacket(make([]byte, len(packet)-1))
require.ErrorIs(t, err, io.ErrShortBuffer)
require.Zero(t, n)
b := make([]byte, len(packet))
n, err = conn.ReadPacket(b)
require.NoError(t, err)
require.Equal(t, packet, b[:n])
n, err = conn.ReadPacket(make([]byte, 1500))
require.NoError(t, err)
require.Equal(t, len(packet), n)
}
func TestCloseConcurrently(t *testing.T) {
for _, side := range []string{"client", "proxy"} {
t.Run(side, func(t *testing.T) {
client, server := setupConns(t)
conn := client
if side == "proxy" {
conn = server
}
readErr := make(chan error, 1)
go func() {
b := make([]byte, 1500)
for {
if _, err := conn.ReadPacket(b); err != nil {
readErr <- err
return
}
}
}()
writeErr := make(chan error, 1)
go func() {
for {
if _, err := conn.WritePacket(ipv4Packet(64, 17, testSrc4, testDst4, nil, nil)); err != nil {
writeErr <- err
return
}
}
}()
var wg sync.WaitGroup
for range 4 {
wg.Go(func() { assert.NoError(t, conn.Close()) })
}
wg.Wait()
for _, errChan := range []chan error{readErr, writeErr} {
select {
case err := <-errChan:
require.ErrorIs(t, err, net.ErrClosed)
case <-time.After(5 * time.Second):
t.Fatal("timeout")
}
}
require.NoError(t, conn.Close())
})
}
}
@@ -0,0 +1,94 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"encoding/binary"
"errors"
"fmt"
"golang.org/x/net/icmp"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
)
const (
ipv4MinMTU = 68
ipv6MinMTU = 1280
)
var ErrMTUTooSmall = errors.New("connect-ip: tunnel MTU below the minimum link MTU")
func composeICMPTooLargePacket(b []byte, mtu int) ([]byte, error) {
if len(b) == 0 {
return nil, errors.New("connect-ip: empty packet")
}
var icmpMessage *icmp.Message
var psh []byte
switch v := ipVersion(b); v {
case 4:
if len(b) < ipv4.HeaderLen {
return nil, errors.New("connect-ip: IPv4 packet too short")
}
if mtu < ipv4MinMTU {
return nil, fmt.Errorf("%w: %d bytes", ErrMTUTooSmall, mtu)
}
icmpMessage = &icmp.Message{
Type: ipv4.ICMPTypeDestinationUnreachable,
Code: 4,
Body: &icmp.PacketTooBig{
MTU: mtu,
Data: b[:min(len(b), max(ipv4.HeaderLen, int(b[0]&0x0f)<<2)+8)],
},
}
case 6:
if len(b) < ipv6.HeaderLen {
return nil, errors.New("connect-ip: IPv6 packet too short")
}
if mtu < ipv6MinMTU {
return nil, fmt.Errorf("%w: %d bytes", ErrMTUTooSmall, mtu)
}
icmpMessage = &icmp.Message{
Type: ipv6.ICMPTypePacketTooBig,
Body: &icmp.PacketTooBig{
MTU: mtu,
Data: b[:min(len(b), 1232)],
},
}
psh = icmp.IPv6PseudoHeader(b[24:40], b[8:24])
default:
return nil, fmt.Errorf("connect-ip: unknown IP version: %d", v)
}
icmp, err := icmpMessage.Marshal(psh)
if err != nil {
return nil, fmt.Errorf("connect-ip: failed to marshal ICMP message: %w", err)
}
if ipVersion(b) == 4 {
var header [ipv4.HeaderLen]byte
header[0] = 4<<4 | ipv4.HeaderLen>>2
ipLen := ipv4.HeaderLen + len(icmp)
binary.BigEndian.PutUint16(header[2:4], uint16(ipLen))
header[8] = 64
header[9] = 1
copy(header[12:16], b[16:20])
copy(header[16:20], b[12:16])
binary.BigEndian.PutUint16(header[10:12], calculateIPv4Checksum(header[:]))
return append(header[:], icmp...), nil
}
var header [ipv6.HeaderLen]byte
header[0] = 6 << 4
binary.BigEndian.PutUint16(header[4:6], uint16(len(icmp)))
header[6] = 58
header[7] = 64
copy(header[8:24], b[24:40])
copy(header[24:40], b[8:24])
return append(header[:], icmp...), nil
}
@@ -0,0 +1,158 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"encoding/binary"
"net"
"net/netip"
"testing"
"github.com/stretchr/testify/require"
"golang.org/x/net/icmp"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
)
func TestICMPTooLargeIPv4(t *testing.T) {
src := netip.MustParseAddr("192.168.1.1")
dst := netip.MustParseAddr("8.8.8.8")
origHdr := &ipv4.Header{
Version: 4,
Len: ipv4.HeaderLen,
TotalLen: 60,
TTL: 64,
Protocol: 6,
Src: src.AsSlice(),
Dst: dst.AsSlice(),
}
origBytes, err := origHdr.Marshal()
require.NoError(t, err)
data, err := composeICMPTooLargePacket(origBytes, 1200)
require.NoError(t, err)
hdr, err := ipv4.ParseHeader(data)
require.NoError(t, err)
require.Equal(t, 4, hdr.Version)
require.Equal(t, ipProtoICMP, hdr.Protocol)
require.Equal(t, dst.String(), hdr.Src.String())
require.Equal(t, src.String(), hdr.Dst.String())
require.Equal(t, uint16(hdr.Checksum), calculateIPv4Checksum(data[:ipv4.HeaderLen]))
icmpMsg, err := icmp.ParseMessage(ipProtoICMP, data[ipv4.HeaderLen:])
require.NoError(t, err)
require.Equal(t, ipv4.ICMPTypeDestinationUnreachable, icmpMsg.Type)
require.Equal(t, 4, icmpMsg.Code)
require.Equal(t, uint16(1200), binary.BigEndian.Uint16(data[ipv4.HeaderLen+6:]))
require.Equal(t, origBytes, data[ipv4.HeaderLen+8:])
}
func TestICMPTooLargeIPv4Options(t *testing.T) {
options := []byte{0x94, 0x04, 0x00, 0x00}
orig := ipv4Packet(64, 6, netip.MustParseAddr("192.168.1.1"), netip.MustParseAddr("8.8.8.8"), options, make([]byte, 20))
data, err := composeICMPTooLargePacket(orig, 1200)
require.NoError(t, err)
require.Equal(t, orig[:ipv4.HeaderLen+len(options)+8], data[ipv4.HeaderLen+8:])
}
func TestICMPTooLargeIPv6(t *testing.T) {
const mtu = 1337
src := netip.MustParseAddr("2001:db8::1")
dst := netip.MustParseAddr("1:2:3:4::5")
orig := []byte{
0x60, 0x00, 0x00, 0x00,
0x00, 0x00,
0x00, 0x2a,
}
orig = append(orig, src.AsSlice()...)
orig = append(orig, dst.AsSlice()...)
orig = append(orig, []byte("foobar")...)
data, err := composeICMPTooLargePacket(orig, mtu)
require.NoError(t, err)
hdr, err := ipv6.ParseHeader(data)
require.NoError(t, err)
require.Equal(t, 6, hdr.Version)
require.Equal(t, ipProtoICMPv6, hdr.NextHeader)
require.Equal(t, dst.String(), hdr.Src.String())
require.Equal(t, src.String(), hdr.Dst.String())
icmpMsg, err := icmp.ParseMessage(ipProtoICMPv6, data[ipv6.HeaderLen:])
require.NoError(t, err)
require.Equal(t, ipv6.ICMPTypePacketTooBig, icmpMsg.Type)
icmpBody, ok := icmpMsg.Body.(*icmp.PacketTooBig)
require.True(t, ok)
require.Equal(t, mtu, icmpBody.MTU)
require.Equal(t, orig, icmpBody.Data)
}
func TestICMPTooLargeMinimumMTU(t *testing.T) {
ipv4Orig := ipv4Packet(64, 6, testSrc4, testDst4, nil, make([]byte, 100))
ipv6Orig := ipv6Packet(64, 6, testSrc6, testDst6, make([]byte, 1300))
for _, tc := range []struct {
name string
packet []byte
mtu int
tooSmall bool
}{
{"IPv4 minimum", ipv4Orig, 68, false},
{"IPv4 below minimum", ipv4Orig, 67, true},
{"IPv4 negative", ipv4Orig, -1, true},
{"IPv6 minimum", ipv6Orig, 1280, false},
{"IPv6 below minimum", ipv6Orig, 1279, true},
} {
t.Run(tc.name, func(t *testing.T) {
data, err := composeICMPTooLargePacket(tc.packet, tc.mtu)
if tc.tooSmall {
require.ErrorIs(t, err, ErrMTUTooSmall)
require.Nil(t, data)
return
}
require.NoError(t, err)
require.NotEmpty(t, data)
})
}
}
func TestICMPFailures(t *testing.T) {
t.Run("empty packet", func(t *testing.T) {
_, err := composeICMPTooLargePacket([]byte{}, 1)
require.EqualError(t, err, "connect-ip: empty packet")
})
t.Run("too short IPv4 header", func(t *testing.T) {
origHdr := &ipv4.Header{
Version: 4,
Len: ipv4.HeaderLen,
TotalLen: 60,
Src: net.IPv4(1, 2, 3, 4),
Dst: net.IPv4(5, 6, 7, 8),
}
data, err := origHdr.Marshal()
require.NoError(t, err)
_, err = composeICMPTooLargePacket(data[:ipv4.HeaderLen-1], 1)
require.EqualError(t, err, "connect-ip: IPv4 packet too short")
})
t.Run("too short IPv6 header", func(t *testing.T) {
data := []byte{
0x60, 0x00, 0x00, 0x00,
0x00, 0x00,
0x00, 0x40,
}
data = append(data, net.ParseIP("2001:db8::1").To16()...)
data = append(data, net.ParseIP("2001:db8::2").To16()...)
_, err := composeICMPTooLargePacket(data[:ipv6.HeaderLen-1], 1)
require.EqualError(t, err, "connect-ip: IPv6 packet too short")
})
t.Run("unknown IP version", func(t *testing.T) {
data := []byte{
0x30, 0x00, 0x00, 0x00,
}
_, err := composeICMPTooLargePacket(data, 1)
require.EqualError(t, err, "connect-ip: unknown IP version: 3")
})
}
@@ -0,0 +1,73 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import "net/netip"
func rangeToPrefixes(start, end netip.Addr) []netip.Prefix {
var prefixes []netip.Prefix
for current := start; current.Compare(end) <= 0; {
prefix := findLargestPrefix(current, end)
prefixes = append(prefixes, prefix)
lastIP := lastIPInPrefix(prefix)
if lastIP.Compare(end) >= 0 {
break
}
current = lastIP.Next()
}
return prefixes
}
func findLargestPrefix(start, end netip.Addr) netip.Prefix {
if start == end {
return netip.PrefixFrom(start, start.BitLen())
}
var prefixLen int
for prefixLen = start.BitLen(); prefixLen > 0; prefixLen-- {
prefix := netip.PrefixFrom(start, prefixLen-1)
if lastIPInPrefix(prefix).Compare(end) > 0 || !isAligned(start, prefixLen-1) {
break
}
}
return netip.PrefixFrom(start, prefixLen)
}
func lastIPInPrefix(prefix netip.Prefix) netip.Addr {
addr := prefix.Addr()
bits := addr.As16()
hostBits := addr.BitLen() - prefix.Bits()
for i := len(bits) - 1; i >= 0 && hostBits > 0; i-- {
bitsInThisByte := min(8, hostBits)
mask := byte((1 << bitsInThisByte) - 1)
bits[i] |= mask
hostBits -= bitsInThisByte
}
if addr.Is4() {
return netip.AddrFrom4([4]byte(bits[12:16]))
}
return netip.AddrFrom16(bits)
}
func isAligned(addr netip.Addr, prefixLen int) bool {
bits := addr.As16()
hostBits := addr.BitLen() - prefixLen
for i := len(bits) - 1; i >= 0 && hostBits > 0; i-- {
bitsInThisByte := min(8, hostBits)
mask := byte((1 << bitsInThisByte) - 1)
if bits[i]&mask != 0 {
return false
}
hostBits -= bitsInThisByte
}
return true
}
@@ -0,0 +1,78 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"fmt"
"net/netip"
"testing"
"github.com/stretchr/testify/require"
)
func TestIPRanges(t *testing.T) {
tests := []struct {
start, end netip.Addr
want []netip.Prefix
}{
{
start: netip.MustParseAddr("192.168.1.1"),
end: netip.MustParseAddr("192.168.1.1"),
want: []netip.Prefix{netip.MustParsePrefix("192.168.1.1/32")},
},
{
start: netip.MustParseAddr("192.168.1.0"),
end: netip.MustParseAddr("192.168.1.1"),
want: []netip.Prefix{netip.MustParsePrefix("192.168.1.0/31")},
},
{
start: netip.MustParseAddr("192.168.1.1"),
end: netip.MustParseAddr("192.168.1.2"),
want: []netip.Prefix{netip.MustParsePrefix("192.168.1.1/32"), netip.MustParsePrefix("192.168.1.2/32")},
},
{
start: netip.MustParseAddr("192.168.1.0"),
end: netip.MustParseAddr("192.168.1.255"),
want: []netip.Prefix{netip.MustParsePrefix("192.168.1.0/24")},
},
{
start: netip.MustParseAddr("10.0.0.0"),
end: netip.MustParseAddr("10.1.0.255"),
want: []netip.Prefix{netip.MustParsePrefix("10.0.0.0/16"), netip.MustParsePrefix("10.1.0.0/24")},
},
{
start: netip.MustParseAddr("2001:0db8:85a3::8a2e:0370:7334"),
end: netip.MustParseAddr("2001:0db8:85a3::8a2e:0370:7334"),
want: []netip.Prefix{netip.MustParsePrefix("2001:0db8:85a3::8a2e:0370:7334/128")},
},
{
start: netip.MustParseAddr("2001:db8::0"),
end: netip.MustParseAddr("2001:db8::ffff:ffff:ffff:ffff"),
want: []netip.Prefix{netip.MustParsePrefix("2001:db8::/64")},
},
{
start: netip.MustParseAddr("2001:db8::1"),
end: netip.MustParseAddr("2001:db8::2"),
want: []netip.Prefix{netip.MustParsePrefix("2001:db8::1/128"), netip.MustParsePrefix("2001:db8::2/128")},
},
{
start: netip.MustParseAddr("2001:db8:1234:5678::"),
end: netip.MustParseAddr("2001:db8:1234:5679::"),
want: []netip.Prefix{
netip.MustParsePrefix("2001:db8:1234:5678::/64"),
netip.MustParsePrefix("2001:db8:1234:5679::/128"),
},
},
}
for _, test := range tests {
t.Run(fmt.Sprintf("%s-%s", test.start, test.end), func(t *testing.T) {
prefixes := rangeToPrefixes(test.start, test.end)
require.Equal(t, test.want, prefixes)
})
}
}
@@ -0,0 +1,30 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"errors"
"net/http"
"github.com/apernet/quic-go/http3"
"github.com/apernet/quic-go/quicvarint"
)
var contextIDZero = quicvarint.Append([]byte{}, 0)
type Proxy struct{}
func (s *Proxy) Proxy(w http.ResponseWriter, _ *ProxyRequest) (*Conn, error) {
streamer, ok := w.(http3.HTTPStreamer)
if !ok {
return nil, errors.New("connect-ip: response writer is not an HTTP/3 stream")
}
w.Header().Set(http3.CapsuleProtocolHeader, capsuleProtocolHeaderValue)
w.WriteHeader(http.StatusOK)
return newProxiedConn(streamer.HTTPStream()), nil
}
@@ -0,0 +1,413 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"context"
"crypto/tls"
"encoding/binary"
"fmt"
"net"
"net/http"
"net/netip"
"slices"
"testing"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
)
func dialHTTP3(t *testing.T, addr string) *http3.ClientConn {
t.Helper()
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
defer cancel()
qconn, err := quic.DialAddr(
ctx,
addr,
&tls.Config{ServerName: "localhost", RootCAs: certPool, NextProtos: []string{http3.NextProtoH3}},
&quic.Config{EnableDatagrams: true, InitialPacketSize: 1350, DisablePathMTUDiscovery: true},
)
require.NoError(t, err)
t.Cleanup(func() { qconn.CloseWithError(0, "") })
return (&http3.Transport{EnableDatagrams: true}).NewClientConn(qconn)
}
func setupConns(t *testing.T) (client, server *Conn) {
t.Helper()
p := &Proxy{}
conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0})
require.NoError(t, err)
t.Cleanup(func() { conn.Close() })
proxyURL := fmt.Sprintf("https://%s/connect-ip", conn.LocalAddr())
connChan := make(chan *Conn, 1)
mux := http.NewServeMux()
mux.HandleFunc("/connect-ip", func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, "Bearer token", r.Header.Get("Authorization"))
mreq, err := ParseProxyRequest(r)
if !assert.NoError(t, err) {
w.WriteHeader(http.StatusBadRequest)
return
}
conn, err := p.Proxy(w, mreq)
if assert.NoError(t, err) {
connChan <- conn
}
})
s := http3.Server{
Handler: mux,
Addr: ":0",
EnableDatagrams: true,
TLSConfig: tlsConf,
}
go func() { s.Serve(conn) }()
t.Cleanup(func() { s.Close() })
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
defer cancel()
req, err := NewRequest(ctx, proxyURL)
require.NoError(t, err)
req.Header().Set("Authorization", "Bearer token")
client, rsp, err := NewClientConn(dialHTTP3(t, conn.LocalAddr().String())).Dial(req)
require.NoError(t, err)
t.Cleanup(func() { client.Close() })
require.Equal(t, http.StatusOK, rsp.StatusCode)
select {
case <-time.After(5 * time.Second):
t.Fatal("timed out")
case server = <-connChan:
}
t.Cleanup(func() { server.Close() })
return client, server
}
func TestAddressAssignment(t *testing.T) {
client, server := setupConns(t)
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond)
defer cancel()
_, err := server.ReceiveAddressAssignment(ctx)
require.ErrorIs(t, err, context.DeadlineExceeded)
ctx, cancel = context.WithTimeout(context.Background(), time.Second)
defer cancel()
pref1 := netip.MustParsePrefix("1.1.1.0/24")
pref2 := netip.MustParsePrefix("2001:db8::/64")
require.NoError(t, client.AssignAddresses([]netip.Prefix{pref1, pref2}))
assigned, err := server.ReceiveAddressAssignment(ctx)
require.NoError(t, err)
require.Equal(t, []AssignedAddress{{IPPrefix: pref1}, {IPPrefix: pref2}}, assigned)
require.NoError(t, client.AssignAddresses([]netip.Prefix{}))
assigned, err = server.ReceiveAddressAssignment(ctx)
require.NoError(t, err)
require.Empty(t, assigned)
}
func TestRejectingAddressRequestKeepsPeerUnrestricted(t *testing.T) {
client, server := setupConns(t)
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
defer cancel()
clientAddr := netip.MustParsePrefix("192.0.2.2/32")
require.NoError(t, server.AssignAddresses([]netip.Prefix{clientAddr}))
_, err := client.ReceiveAddressAssignment(ctx)
require.NoError(t, err)
_, err = server.RequestAddresses([]netip.Prefix{netip.MustParsePrefix("0.0.0.0/32")})
require.NoError(t, err)
req, err := client.ReceiveAddressRequest(ctx)
require.NoError(t, err)
require.NoError(t, req.Respond([]netip.Prefix{{}}, nil))
assigned, err := server.ReceiveAddressAssignment(ctx)
require.NoError(t, err)
require.Len(t, assigned, 1)
require.True(t, assigned[0].Rejected())
packet := ipv4Packet(64, 17, netip.MustParseAddr("203.0.113.9"), clientAddr.Addr(), nil, []byte("foobar"))
_, err = server.WritePacket(slices.Clone(packet))
require.NoError(t, err)
received := make(chan []byte, 1)
go func() {
b := make([]byte, 1500)
if n, err := client.ReadPacket(b); err == nil {
received <- b[:n]
}
}()
select {
case b := <-received:
require.Equal(t, packet[20:], b[20:])
case <-ctx.Done():
t.Fatal("packet was not received")
}
}
func TestRejectingAddressRequestWithdrawsAssignment(t *testing.T) {
client, server := setupConns(t)
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
defer cancel()
clientAddr := netip.MustParsePrefix("192.0.2.2/32")
dst := netip.MustParseAddr("198.51.100.1")
require.NoError(t, server.AssignAddresses([]netip.Prefix{clientAddr}))
require.NoError(t, server.AdvertiseRoute([]IPRoute{{StartIP: dst, EndIP: dst}}))
_, err := client.ReceiveAddressAssignment(ctx)
require.NoError(t, err)
received := make(chan []byte, 4)
go func() {
b := make([]byte, 1500)
for {
n, err := server.ReadPacket(b)
if err != nil {
return
}
received <- slices.Clone(b[20:n])
}
}()
send := func(payload string) {
_, err := client.WritePacket(ipv4Packet(64, 17, clientAddr.Addr(), dst, nil, []byte(payload)))
require.NoError(t, err)
}
send("assigned")
select {
case b := <-received:
require.Equal(t, "assigned", string(b))
case <-ctx.Done():
t.Fatal("packet from the assigned address was not received")
}
_, err = client.RequestAddresses([]netip.Prefix{netip.MustParsePrefix("0.0.0.0/32")})
require.NoError(t, err)
req, err := server.ReceiveAddressRequest(ctx)
require.NoError(t, err)
require.NoError(t, req.Respond([]netip.Prefix{{}}, nil))
assigned, err := client.ReceiveAddressAssignment(ctx)
require.NoError(t, err)
require.Len(t, assigned, 1)
require.True(t, assigned[0].Rejected())
send("withdrawn")
select {
case b := <-received:
t.Fatalf("packet from a withdrawn address was received: %q", b)
case <-time.After(200 * time.Millisecond):
}
}
func TestRouteAdvertisement(t *testing.T) {
client, server := setupConns(t)
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond)
defer cancel()
_, err := server.Routes(ctx)
require.ErrorIs(t, err, context.DeadlineExceeded)
ctx, cancel = context.WithTimeout(context.Background(), time.Second)
defer cancel()
require.ErrorContains(t,
client.AdvertiseRoute([]IPRoute{
{StartIP: netip.MustParseAddr("1.1.1.2"), EndIP: netip.MustParseAddr("1.1.1.1"), IPProtocol: 42},
}),
"connect-ip: invalid route 0: start IP 1.1.1.2 is greater than end IP 1.1.1.1",
)
require.NoError(t, client.AdvertiseRoute([]IPRoute{
{StartIP: netip.MustParseAddr("1.1.1.1"), EndIP: netip.MustParseAddr("2.2.2.2"), IPProtocol: 42},
{StartIP: netip.MustParseAddr("2001:db8::1"), EndIP: netip.MustParseAddr("2001:db8::100"), IPProtocol: 24},
}))
routes, err := server.Routes(ctx)
require.NoError(t, err)
require.Equal(t, []IPRoute{
{StartIP: netip.MustParseAddr("1.1.1.1"), EndIP: netip.MustParseAddr("2.2.2.2"), IPProtocol: 42},
{StartIP: netip.MustParseAddr("2001:db8::1"), EndIP: netip.MustParseAddr("2001:db8::100"), IPProtocol: 24},
}, routes)
require.NoError(t, client.AdvertiseRoute([]IPRoute{}))
routes, err = server.Routes(ctx)
require.NoError(t, err)
require.Empty(t, routes)
}
func TestTTLs(t *testing.T) {
t.Run("IPv4", func(t *testing.T) {
client, server := setupConns(t)
require.NoError(t, server.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("192.168.1.1/32")}))
require.NoError(t, server.AdvertiseRoute([]IPRoute{
{StartIP: netip.MustParseAddr("0.0.0.0"), EndIP: netip.MustParseAddr("255.255.255.255")},
}))
src, dst := netip.MustParseAddr("192.168.1.1"), netip.MustParseAddr("8.8.8.8")
icmp, err := client.WritePacket(ipv4Packet(1, 0, src, dst, nil, nil))
require.NoError(t, err)
require.Empty(t, icmp)
icmp, err = client.WritePacket(ipv4Packet(42, 0, src, dst, nil, nil))
require.NoError(t, err)
require.Empty(t, icmp)
receivedPacket := make([]byte, 1500)
n, err := server.ReadPacket(receivedPacket)
require.NoError(t, err)
receivedPacket = receivedPacket[:n]
receivedHdr, err := ipv4.ParseHeader(receivedPacket)
require.NoError(t, err)
require.Equal(t, uint16(receivedHdr.Checksum), calculateIPv4Checksum(receivedPacket[:ipv4.HeaderLen]))
require.Equal(t, 41, receivedHdr.TTL)
})
t.Run("IPv6", func(t *testing.T) {
client, server := setupConns(t)
require.NoError(t, server.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("2001:db8::1/128")}))
require.NoError(t, server.AdvertiseRoute([]IPRoute{
{StartIP: netip.MustParseAddr("::"), EndIP: netip.MustParseAddr("ffff:ffff:ffff:ffff:ffff:ffff:ffff:ffff")},
}))
packetHopLimit1 := []byte{
0x60, 0x00, 0x00, 0x00,
0x00, 0x00,
0x00, 0x01,
0x20, 0x01, 0x0d, 0xb8, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01,
0x20, 0x01, 0x48, 0x60, 0x48, 0x60, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x88, 0x88,
}
icmp, err := client.WritePacket(packetHopLimit1)
require.NoError(t, err)
require.Empty(t, icmp)
packet := []byte{
0x60, 0x00, 0x00, 0x00,
0x00, 0x00,
0x00, 0x2A,
0x20, 0x01, 0x0d, 0xb8, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01,
0x20, 0x01, 0x48, 0x60, 0x48, 0x60, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x88, 0x88,
}
icmp, err = client.WritePacket(packet)
require.NoError(t, err)
require.Empty(t, icmp)
receivedPacket := make([]byte, 1500)
n, err := server.ReadPacket(receivedPacket)
require.NoError(t, err)
receivedPacket = receivedPacket[:n]
receivedHdr, err := ipv6.ParseHeader(receivedPacket)
require.NoError(t, err)
require.Equal(t, 41, receivedHdr.HopLimit)
})
}
func TestMaxPacketSizeOverQUIC(t *testing.T) {
client, server := setupConns(t)
require.NoError(t, server.AdvertiseRoute([]IPRoute{
{StartIP: netip.MustParseAddr("0.0.0.0"), EndIP: netip.MustParseAddr("255.255.255.255")},
}))
size := client.MaxPacketSize()
require.Greater(t, size, 1200)
require.Less(t, size, 1350)
icmp, err := client.WritePacket(ipv4Packet(64, 17, testSrc4, testDst4, nil, make([]byte, size-ipv4.HeaderLen)))
require.NoError(t, err)
require.Nil(t, icmp)
type readResult struct {
n int
err error
}
received := make(chan readResult, 1)
go func() {
n, err := server.ReadPacket(make([]byte, 1500))
received <- readResult{n, err}
}()
select {
case r := <-received:
require.NoError(t, r.err)
require.Equal(t, size, r.n)
case <-time.After(5 * time.Second):
t.Fatal("timeout")
}
icmp, err = client.WritePacket(ipv4Packet(64, 17, testSrc4, testDst4, nil, make([]byte, size+1-ipv4.HeaderLen)))
require.NoError(t, err)
require.NotNil(t, icmp)
require.Equal(t, uint16(size), binary.BigEndian.Uint16(icmp[ipv4.HeaderLen+6:]))
}
func TestClosing(t *testing.T) {
ipv6Packet := []byte{
0x60, 0x00, 0x00, 0x00,
0x00, 0x00,
0x00, 0x2A,
0x20, 0x01, 0x0d, 0xb8, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01,
0x20, 0x01, 0x48, 0x60, 0x48, 0x60, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x88, 0x88,
}
client, server := setupConns(t)
routeErrChan := make(chan error, 1)
prefixErrChan := make(chan error, 1)
go func() {
_, err := server.Routes(context.Background())
routeErrChan <- err
}()
go func() {
_, err := server.ReceiveAddressAssignment(context.Background())
prefixErrChan <- err
}()
require.NoError(t, client.Close())
_, err := client.ReceiveAddressAssignment(context.Background())
require.ErrorIs(t, err, net.ErrClosed)
var closeErr *CloseError
require.ErrorAs(t, err, &closeErr)
require.False(t, closeErr.Remote)
_, err = client.Routes(context.Background())
require.ErrorIs(t, err, net.ErrClosed)
require.ErrorIs(t,
client.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("1.1.1.0/24")}),
net.ErrClosed,
)
require.ErrorIs(t,
client.AdvertiseRoute([]IPRoute{
{StartIP: netip.MustParseAddr("1.1.1.0"), EndIP: netip.MustParseAddr("1.1.1.1"), IPProtocol: 42},
}),
net.ErrClosed,
)
_, err = client.ReadPacket([]byte{0})
require.ErrorIs(t, err, net.ErrClosed)
_, err = client.WritePacket(ipv6Packet)
require.ErrorIs(t, err, net.ErrClosed)
select {
case err := <-routeErrChan:
require.ErrorIs(t, err, net.ErrClosed)
case <-time.After(time.Second):
t.Fatal("timeout")
}
select {
case err := <-prefixErrChan:
require.ErrorIs(t, err, net.ErrClosed)
case <-time.After(time.Second):
t.Fatal("timeout")
}
_, err = server.ReadPacket([]byte{0})
require.ErrorIs(t, err, net.ErrClosed)
_, err = server.WritePacket(ipv6Packet)
require.ErrorIs(t, err, net.ErrClosed)
}
@@ -0,0 +1,91 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"context"
"errors"
"fmt"
"net/http"
"strings"
"github.com/apernet/quic-go/http3"
)
const requestProtocol = "connect-ip"
const capsuleProtocolHeaderValue = "?1"
type Request struct {
req *http.Request
}
func NewRequest(ctx context.Context, rawURL string) (*Request, error) {
if strings.ContainsAny(rawURL, "{}") {
return nil, errors.New("connect-ip: IP flow forwarding not supported: URL contains a URI Template expression")
}
req, err := http.NewRequestWithContext(ctx, http.MethodConnect, rawURL, nil)
if err != nil {
return nil, fmt.Errorf("connect-ip: failed to create request: %w", err)
}
if req.URL.Scheme != "https" || req.URL.Host == "" || !strings.HasPrefix(req.URL.Path, "/") {
return nil, fmt.Errorf("connect-ip: invalid proxy URL %q: expected an absolute https URL with a host and a path", rawURL)
}
req.Proto = requestProtocol
req.Host = req.URL.Host
req.Header.Set(http3.CapsuleProtocolHeader, capsuleProtocolHeaderValue)
return &Request{req: req}, nil
}
func (r *Request) Header() http.Header { return r.req.Header }
func (r *Request) httpRequest() *http.Request { return r.req }
type ProxyRequest struct{}
type ProxyRequestParseError struct {
HTTPStatus int
Err error
}
func (e *ProxyRequestParseError) Error() string { return e.Err.Error() }
func (e *ProxyRequestParseError) Unwrap() error { return e.Err }
func ParseProxyRequest(r *http.Request) (*ProxyRequest, error) {
if r.Method != http.MethodConnect {
return nil, &ProxyRequestParseError{
HTTPStatus: http.StatusMethodNotAllowed,
Err: fmt.Errorf("expected CONNECT request, got %s", r.Method),
}
}
if r.Proto != requestProtocol {
return nil, &ProxyRequestParseError{
HTTPStatus: http.StatusNotImplemented,
Err: fmt.Errorf("unexpected protocol: %s", r.Proto),
}
}
capsuleHeaderValues, ok := r.Header[http3.CapsuleProtocolHeader]
if !ok {
return nil, &ProxyRequestParseError{
HTTPStatus: http.StatusBadRequest,
Err: fmt.Errorf("missing Capsule-Protocol header"),
}
}
if !isCapsuleProtocolEnabled(capsuleHeaderValues) {
return nil, &ProxyRequestParseError{
HTTPStatus: http.StatusBadRequest,
Err: fmt.Errorf("invalid capsule header value: %s", capsuleHeaderValues),
}
}
return &ProxyRequest{}, nil
}
func isCapsuleProtocolEnabled(values []string) bool {
v := strings.Trim(strings.Join(values, ","), " ")
return v == capsuleProtocolHeaderValue || strings.HasPrefix(v, capsuleProtocolHeaderValue+";")
}
@@ -0,0 +1,119 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/apernet/quic-go/http3"
"github.com/stretchr/testify/require"
)
func newRequest(target string) *http.Request {
req := httptest.NewRequest(http.MethodGet, target, nil)
req.Method = http.MethodConnect
req.Proto = requestProtocol
req.Header.Add("Capsule-Protocol", capsuleProtocolHeaderValue)
return req
}
func TestNewRequest(t *testing.T) {
req, err := NewRequest(t.Context(), "https://localhost:1234/masque/ip")
require.NoError(t, err)
httpReq := req.httpRequest()
require.Equal(t, http.MethodConnect, httpReq.Method)
require.Equal(t, requestProtocol, httpReq.Proto)
require.Equal(t, "localhost:1234", httpReq.Host)
require.Equal(t, "?1", req.Header().Get(http3.CapsuleProtocolHeader))
req.Header().Set("Authorization", "Bearer token")
require.Equal(t, "Bearer token", httpReq.Header.Get("Authorization"))
}
func TestNewRequestInvalidURL(t *testing.T) {
for _, tc := range []struct {
name, url, err string
}{
{"template with variables", "https://localhost/.well-known/masque/ip/{target}/{ipproto}/", "IP flow forwarding not supported"},
{"template with query variables", "https://localhost/masque/ip{?target,ipproto}", "IP flow forwarding not supported"},
{"not https", "http://localhost/masque/ip", "expected an absolute https URL"},
{"no host", "https:///masque/ip", "expected an absolute https URL"},
{"no path", "https://localhost", "expected an absolute https URL"},
{"relative", "/masque/ip", "expected an absolute https URL"},
{"unparsable", "https://local\x7fhost/", "failed to create request"},
} {
t.Run(tc.name, func(t *testing.T) {
_, err := NewRequest(t.Context(), tc.url)
require.ErrorContains(t, err, tc.err)
})
}
}
func TestProxyRequestParsing(t *testing.T) {
t.Run("valid request", func(t *testing.T) {
req := newRequest("https://localhost:1234/masque/ip")
r, err := ParseProxyRequest(req)
require.NoError(t, err)
require.Equal(t, &ProxyRequest{}, r)
})
t.Run("wrong protocol", func(t *testing.T) {
req := newRequest("https://localhost:1234/masque")
req.Proto = "not-connect-ip"
_, err := ParseProxyRequest(req)
require.EqualError(t, err, "unexpected protocol: not-connect-ip")
require.Equal(t, http.StatusNotImplemented, err.(*ProxyRequestParseError).HTTPStatus)
})
t.Run("wrong request method", func(t *testing.T) {
req := newRequest("https://localhost:1234/masque")
req.Method = http.MethodHead
_, err := ParseProxyRequest(req)
require.EqualError(t, err, "expected CONNECT request, got HEAD")
require.Equal(t, http.StatusMethodNotAllowed, err.(*ProxyRequestParseError).HTTPStatus)
})
t.Run("missing Capsule-Protocol header", func(t *testing.T) {
req := newRequest("https://localhost:1234/masque")
req.Header.Del("Capsule-Protocol")
_, err := ParseProxyRequest(req)
require.EqualError(t, err, "missing Capsule-Protocol header")
require.Equal(t, http.StatusBadRequest, err.(*ProxyRequestParseError).HTTPStatus)
})
for _, tc := range []struct {
name string
values []string
valid bool
}{
{name: "true", values: []string{"?1"}, valid: true},
{name: "surrounding spaces", values: []string{" ?1 "}, valid: true},
{name: "parameters", values: []string{"?1;a;b=?0;c=\"x\""}, valid: true},
{name: "false", values: []string{"?0"}},
{name: "integer", values: []string{"1"}},
{name: "empty", values: []string{""}},
{name: "not a structured field", values: []string{"🤡"}},
{name: "longer token", values: []string{"?10"}},
{name: "space before parameters", values: []string{"?1 ;a"}},
{name: "list", values: []string{"?1, ?1"}},
{name: "multiple field lines", values: []string{"?1", "?1"}},
} {
t.Run("Capsule-Protocol header: "+tc.name, func(t *testing.T) {
req := newRequest("https://localhost:1234/masque")
req.Header[http3.CapsuleProtocolHeader] = tc.values
_, err := ParseProxyRequest(req)
if tc.valid {
require.NoError(t, err)
return
}
require.ErrorContains(t, err, "invalid capsule header value")
require.Equal(t, http.StatusBadRequest, err.(*ProxyRequestParseError).HTTPStatus)
})
}
}
+237
View File
@@ -0,0 +1,237 @@
package masque
import (
"context"
"net/netip"
"reflect"
"runtime"
"strings"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/net/cnc"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/finalmask"
"github.com/xtls/xray-core/transport/internet/hysteria/congestion"
"github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr"
"github.com/xtls/xray-core/transport/internet/masque/connectip"
"github.com/xtls/xray-core/transport/internet/stat"
"github.com/xtls/xray-core/transport/internet/tls"
)
const (
MinPacketSize = 1280
initialPacketSize = 1350
)
func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (stat.Connection, error) {
tlsConfig := tls.ConfigFromStreamSettings(streamSettings)
if tlsConfig == nil {
return nil, errors.New("tls config is nil")
}
config := streamSettings.ProtocolSettings.(*Config)
dest.Network = net.Network_UDP
gotlsConfig := tlsConfig.GetTLSConfig(tls.WithDestination(dest))
gotlsConfig.NextProtos = []string{http3.NextProtoH3}
quicParams := streamSettings.QuicParams
if quicParams == nil {
quicParams = &internet.QuicParams{
BbrProfile: string(bbr.ProfileStandard),
}
}
quicConfig := &quic.Config{
InitialStreamReceiveWindow: quicParams.InitStreamReceiveWindow,
MaxStreamReceiveWindow: quicParams.MaxStreamReceiveWindow,
InitialConnectionReceiveWindow: quicParams.InitConnReceiveWindow,
MaxConnectionReceiveWindow: quicParams.MaxConnReceiveWindow,
MaxIdleTimeout: time.Duration(quicParams.MaxIdleTimeout) * time.Second,
KeepAlivePeriod: time.Duration(quicParams.KeepAlivePeriod) * time.Second,
MaxIncomingStreams: -1,
InitialPacketSize: initialPacketSize,
DisablePathMTUDiscovery: quicParams.DisablePathMtuDiscovery || (runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin"),
EnableDatagrams: true,
DisablePathManager: true,
}
if quicParams.MaxIdleTimeout == 0 {
quicConfig.MaxIdleTimeout = 30 * time.Second
}
if quicParams.KeepAlivePeriod == 0 {
quicConfig.KeepAlivePeriod = net.QuicgoH3KeepAlivePeriod
}
var pktConn net.PacketConn
var udpAddr net.Addr
if streamSettings.FinalMask != nil {
conn, err := streamSettings.FinalMask.DialUDP(ctx, dest)
if err != nil {
return nil, errors.New("failed to dial to dest").Base(err)
}
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
udpAddr = conn.RemoteAddr()
} else {
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
if err != nil {
return nil, errors.New("failed to dial to dest").Base(err)
}
switch c := conn.(type) {
case *internet.PacketConnWrapper:
pktConn = c.PacketConn
udpAddr = c.RemoteAddr()
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
default:
panic(reflect.TypeOf(c))
}
}
tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO}
qconn, err := tr.Dial(ctx, udpAddr, gotlsConfig, quicConfig)
if err != nil {
tr.Close()
pktConn.Close()
return nil, err
}
context.AfterFunc(qconn.Context(), func() { tr.Close(); pktConn.Close() })
switch quicParams.Congestion {
case "reno":
case "", "bbr", "brutal":
congestion.UseBBR(qconn, bbr.Profile(quicParams.BbrProfile))
case "force-brutal":
congestion.UseBrutal(qconn, quicParams.BrutalUp, quicParams.BrutalDisableLossCompensation)
default:
qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "")
return nil, errors.New("unknown congestion control: ", quicParams.Congestion)
}
conn, err := establish(ctx, qconn, config, authority(config, gotlsConfig.ServerName, dest.Port))
if err != nil {
qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "")
return nil, err
}
return conn, nil
}
func establish(ctx context.Context, qconn *quic.Conn, config *Config, host string) (*Conn, error) {
stop := context.AfterFunc(ctx, func() {
qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeRequestCanceled), "")
})
defer stop()
req, err := connectip.NewRequest(ctx, "https://"+host+config.Path)
if err != nil {
return nil, err
}
header := req.Header()
for k, v := range config.Headers {
header.Set(k, v)
}
switch header.Get("User-Agent") {
case "":
header["User-Agent"] = nil
case "chrome":
header.Set("User-Agent", utils.ChromeUA)
case "firefox":
header.Set("User-Agent", utils.FirefoxUA)
case "safari":
header.Set("User-Agent", utils.SafariUA)
case "edge":
header.Set("User-Agent", utils.MSEdgeUA)
case "curl":
header.Set("User-Agent", utils.CurlUA)
case "golang":
header.Del("User-Agent")
}
cc := (&http3.Transport{EnableDatagrams: true, DisableCompression: true}).NewClientConn(qconn)
ipConn, _, err := connectip.NewClientConn(cc).Dial(req)
if err != nil {
if ctx.Err() != nil {
err = context.Cause(ctx)
}
return nil, errors.New("CONNECT-IP request failed").Base(err)
}
if n := ipConn.MaxPacketSize(); n < MinPacketSize {
ipConn.Close()
return nil, errors.New("the tunnel can only carry ", n, "-byte packets, less than ", MinPacketSize)
}
if _, err := ipConn.RequestAddresses([]netip.Prefix{
netip.PrefixFrom(netip.IPv4Unspecified(), 32),
netip.PrefixFrom(netip.IPv6Unspecified(), 128),
}); err != nil {
ipConn.Close()
return nil, err
}
var local []netip.Addr
for len(local) == 0 {
assigned, err := ipConn.ReceiveAddressAssignment(ctx)
if err != nil {
ipConn.Close()
return nil, errors.New("no address assigned").Base(err)
}
local = localAddrs(assigned)
}
if !stop() {
ipConn.Close()
return nil, errors.New("no address assigned").Base(context.Cause(ctx))
}
conn := &Conn{
ipConn: ipConn,
quicConn: qconn,
local: local,
}
go conn.serveAddressAssignments()
go conn.serveAddressRequests()
return conn, nil
}
func localAddrs(assigned []connectip.AssignedAddress) []netip.Addr {
var local []netip.Addr
var has4, has6 bool
for _, a := range assigned {
if a.Rejected() {
continue
}
addr := a.IPPrefix.Addr()
if a.IPPrefix.Bits() != addr.BitLen() {
addr = a.IPPrefix.Masked().Addr().Next()
}
if addr.Is4() && !has4 {
has4 = true
local = append(local, addr)
} else if addr.Is6() && !has6 {
has6 = true
local = append(local, addr)
}
}
return local
}
func authority(config *Config, serverName string, port net.Port) string {
if config.Host != "" {
return config.Host
}
host := strings.TrimSuffix(strings.TrimPrefix(serverName, "["), "]")
if port == 443 {
if addr, err := netip.ParseAddr(host); err == nil && addr.Is6() {
return "[" + host + "]"
}
return host
}
return net.JoinHostPort(host, port.String())
}
func init() {
common.Must(internet.RegisterTransportDialer(protocolName, Dial))
}
+60
View File
@@ -0,0 +1,60 @@
package masque
import (
"net/netip"
"slices"
"testing"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/masque/connectip"
)
func TestAuthority(t *testing.T) {
for _, c := range []struct {
host, serverName string
port net.Port
want string
}{
{serverName: "example.com", port: 443, want: "example.com"},
{serverName: "example.com", port: 8443, want: "example.com:8443"},
{serverName: "127.0.0.1", port: 443, want: "127.0.0.1"},
{serverName: "[2001:db8::1]", port: 443, want: "[2001:db8::1]"},
{serverName: "[2001:db8::1]", port: 8443, want: "[2001:db8::1]:8443"},
{serverName: "2001:db8::1", port: 8443, want: "[2001:db8::1]:8443"},
{host: "proxy.example", serverName: "example.com", port: 8443, want: "proxy.example"},
} {
if got := authority(&Config{Host: c.host}, c.serverName, c.port); got != c.want {
t.Errorf("authority(%q, %q, %d) = %q, want %q", c.host, c.serverName, c.port, got, c.want)
}
}
}
func TestLocalAddrs(t *testing.T) {
assigned := func(prefixes ...string) []connectip.AssignedAddress {
var a []connectip.AssignedAddress
for _, p := range prefixes {
a = append(a, connectip.AssignedAddress{IPPrefix: netip.MustParsePrefix(p)})
}
return a
}
addrs := func(s ...string) []netip.Addr {
var a []netip.Addr
for _, v := range s {
a = append(a, netip.MustParseAddr(v))
}
return a
}
for _, c := range []struct {
assigned []connectip.AssignedAddress
want []netip.Addr
}{
{assigned("192.0.2.2/32", "2001:db8::2/128"), addrs("192.0.2.2", "2001:db8::2")},
{assigned("2001:db8::/64", "192.0.2.0/24", "198.51.100.7/32"), addrs("2001:db8::1", "192.0.2.1")},
{assigned("0.0.0.0/32", "2001:db8::2/128"), addrs("2001:db8::2")},
{assigned("0.0.0.0/32", "::/128"), nil},
} {
if got := localAddrs(c.assigned); !slices.Equal(got, c.want) {
t.Errorf("localAddrs(%v) = %v, want %v", c.assigned, got, c.want)
}
}
}