Remove Snell and OpenVPN

This commit is contained in:
Shtorm
2026-09-02 21:01:21 +03:00
parent e8f6936480
commit 99bbd49406
51 changed files with 5 additions and 6394 deletions
-20
View File
@@ -23,11 +23,9 @@ builds:
- with_tailscale
- with_masque
- with_mtproxy
- with_openvpn
- with_trusttunnel
- with_call
- with_sudoku
- with_snell
- with_manager
- with_admin_panel
- with_profiler
@@ -63,11 +61,9 @@ builds:
- with_tailscale
- with_masque
- with_mtproxy
- with_openvpn
- with_trusttunnel
- with_call
- with_sudoku
- with_snell
- with_profiler
- badlinkname
- tfogo_checklinkname0
@@ -124,11 +120,9 @@ builds:
- with_tailscale
- with_masque
- with_mtproxy
- with_openvpn
- with_trusttunnel
- with_call
- with_sudoku
- with_snell
- with_manager
- with_admin_panel
- with_profiler
@@ -159,11 +153,9 @@ builds:
- with_tailscale
- with_masque
- with_mtproxy
- with_openvpn
- with_trusttunnel
- with_call
- with_sudoku
- with_snell
- with_manager
- with_admin_panel
- with_profiler
@@ -194,11 +186,9 @@ builds:
- with_tailscale
- with_masque
- with_mtproxy
- with_openvpn
- with_trusttunnel
- with_call
- with_sudoku
- with_snell
- with_manager
- with_admin_panel
- with_profiler
@@ -229,11 +219,9 @@ builds:
- with_tailscale
- with_masque
- with_mtproxy
- with_openvpn
- with_trusttunnel
- with_call
- with_sudoku
- with_snell
- with_manager
- with_admin_panel
- with_profiler
@@ -264,11 +252,9 @@ builds:
- with_tailscale
- with_masque
- with_mtproxy
- with_openvpn
- with_trusttunnel
- with_call
- with_sudoku
- with_snell
- with_manager
- with_admin_panel
- with_profiler
@@ -315,11 +301,9 @@ builds:
- with_tailscale
- with_masque
- with_mtproxy
- with_openvpn
- with_trusttunnel
- with_call
- with_sudoku
- with_snell
- with_manager
- with_admin_panel
- with_profiler
@@ -374,11 +358,9 @@ builds:
- with_tailscale
- with_masque
- with_mtproxy
- with_openvpn
- with_trusttunnel
- with_call
- with_sudoku
- with_snell
- with_profiler
- badlinkname
- tfogo_checklinkname0
@@ -448,11 +430,9 @@ builds:
- with_mtproxy
- with_ccm
- with_ocm
- with_openvpn
- with_trusttunnel
- with_call
- with_sudoku
- with_snell
- with_profiler
- badlinkname
- tfogo_checklinkname0
-2
View File
@@ -14,10 +14,8 @@ Sing-box with extended features.
- **MASQUE** — Cloudflare MASQUE proxy over QUIC / HTTP-2
- **MTProxy** — Telegram MTProxy server with FakeTLS and domain fronting
- **Mieru** — Secure, hard to classify, hard to probe network protocol
- **OpenVPN** — OpenVPN client with tls-auth, tls-crypt and tls-crypt-v2 support
- **TrustTunnel** — AdGuard's obfuscated VPN protocol, indistinguishable from HTTPS traffic
- **Sudoku** — Traffic obfuscation protocol based on 4×4 Sudoku puzzles with low-entropy fingerprints
- **Snell** — Lightweight encrypted proxy (v1v5) with TLS / HTTP obfuscation
- **SSH** — SSH client and server with certificate authentication and upstream fallback
- **Call** — Traffic tunneling through video-call platforms (VK, Dion, Telemost, WBStream)
- **VPN** — Routed tunnel over any sing-box protocol
+1 -1
View File
@@ -63,7 +63,7 @@ func init() {
sharedFlags = append(sharedFlags, "-ldflags", "-X github.com/sagernet/sing-box/constant.Version="+currentTag+" -X internal/godebug.defaultGODEBUG=multipathtcp=0 -s -w -buildid= -checklinkname=0")
debugFlags = append(debugFlags, "-ldflags", "-X github.com/sagernet/sing-box/constant.Version="+currentTag+" -X internal/godebug.defaultGODEBUG=multipathtcp=0 -checklinkname=0")
sharedTags = append(sharedTags, "with_gvisor", "with_quic", "with_wireguard", "with_masque", "with_mtproxy", "with_trusttunnel", "with_call", "with_openvpn", "with_sudoku", "with_snell", "with_utls", "with_naive_outbound", "with_clash_api", "badlinkname", "tfogo_checklinkname0")
sharedTags = append(sharedTags, "with_gvisor", "with_quic", "with_wireguard", "with_masque", "with_mtproxy", "with_trusttunnel", "with_call", "with_sudoku", "with_utls", "with_naive_outbound", "with_clash_api", "badlinkname", "tfogo_checklinkname0")
darwinTags = append(darwinTags, "with_dhcp", "grpcnotrace")
// memcTags = append(memcTags, "with_tailscale")
sharedTags = append(sharedTags, "with_tailscale", "ts_omit_logtail", "ts_omit_ssh", "ts_omit_drive", "ts_omit_taildrop", "ts_omit_webclient", "ts_omit_doctor", "ts_omit_capture", "ts_omit_kube", "ts_omit_aws", "ts_omit_synology", "ts_omit_bird")
-135
View File
@@ -1,135 +0,0 @@
package tls
import (
"context"
"crypto/tls"
"crypto/x509"
"encoding/pem"
"os"
"strings"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/option"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
)
func NewOpenVPNClient(ctx context.Context, logger logger.ContextLogger, options option.OpenVPNTLSOptions) (Config, error) {
ca := options.CA
if ca == "" && options.CAPath != "" {
data, err := os.ReadFile(options.CAPath)
if err != nil {
return nil, E.Cause(err, "read ca_path")
}
ca = string(data)
}
certificate := options.Certificate
if certificate == "" && options.CertificatePath != "" {
data, err := os.ReadFile(options.CertificatePath)
if err != nil {
return nil, E.Cause(err, "read certificate_path")
}
certificate = string(data)
}
key := options.Key
if key == "" && options.KeyPath != "" {
data, err := os.ReadFile(options.KeyPath)
if err != nil {
return nil, E.Cause(err, "read key_path")
}
key = string(data)
}
if strings.TrimSpace(ca) == "" {
return nil, E.New("openvpn: missing ca certificate")
}
if block, _ := pem.Decode([]byte(ca)); block == nil {
return nil, E.New("openvpn: ca is not valid PEM")
}
hasCert := strings.TrimSpace(certificate) != "" || strings.TrimSpace(key) != ""
if hasCert {
if strings.TrimSpace(certificate) == "" || strings.TrimSpace(key) == "" {
return nil, E.New("openvpn: certificate and key must both be set")
}
if block, _ := pem.Decode([]byte(certificate)); block == nil {
return nil, E.New("openvpn: certificate is not valid PEM")
}
if block, _ := pem.Decode([]byte(key)); block == nil {
return nil, E.New("openvpn: key is not valid PEM")
}
}
roots := x509.NewCertPool()
if !roots.AppendCertsFromPEM([]byte(ca)) {
return nil, E.New("openvpn: failed to parse ca certificate")
}
var tlsConfig tls.Config
tlsConfig.RootCAs = roots
tlsConfig.InsecureSkipVerify = true
if options.CipherSuites != nil {
find:
for _, cipherSuite := range options.CipherSuites {
for _, tlsCipherSuite := range tls.CipherSuites() {
if cipherSuite == tlsCipherSuite.Name {
tlsConfig.CipherSuites = append(tlsConfig.CipherSuites, tlsCipherSuite.ID)
continue find
}
}
return nil, E.New("unknown cipher_suite: ", cipherSuite)
}
}
tlsConfig.VerifyConnection = func(cs tls.ConnectionState) error {
if len(cs.PeerCertificates) == 0 {
return E.New("openvpn: server did not provide certificate")
}
cert := cs.PeerCertificates[0]
intermediates := x509.NewCertPool()
for _, intermediate := range cs.PeerCertificates[1:] {
intermediates.AddCert(intermediate)
}
_, err := cert.Verify(x509.VerifyOptions{
Roots: roots,
Intermediates: intermediates,
KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
})
if err != nil {
return err
}
if options.VerifyX509Name != "" {
cn := cert.Subject.CommonName
switch options.VerifyX509NameMode {
case "name-prefix":
if !strings.HasPrefix(cn, options.VerifyX509Name) {
return E.New("openvpn: server CN ", cn, " does not match prefix ", options.VerifyX509Name)
}
case "name-suffix":
if !strings.HasSuffix(cn, options.VerifyX509Name) {
return E.New("openvpn: server CN ", cn, " does not match suffix ", options.VerifyX509Name)
}
default:
if cn != options.VerifyX509Name {
return E.New("openvpn: server CN ", cn, " does not match ", options.VerifyX509Name)
}
}
}
return nil
}
if hasCert {
cert, err := tls.X509KeyPair([]byte(certificate), []byte(key))
if err != nil {
return nil, E.Cause(err, "openvpn: parse client certificate/key")
}
tlsConfig.Certificates = []tls.Certificate{cert}
}
var config Config = &STDClientConfig{ctx, &tlsConfig, false, 0, false}
if options.KernelRx || options.KernelTx {
if !C.IsLinux {
return nil, E.New("kTLS is only supported on Linux")
}
config = &KTLSClientConfig{
Config: config,
logger: logger,
kernelTx: options.KernelTx,
kernelRx: options.KernelRx,
}
}
return config, nil
}
-6
View File
@@ -18,7 +18,6 @@ const (
TypeWireGuard = "wireguard"
TypeWARP = "warp"
TypeMASQUE = "masque"
TypeOpenVPN = "openvpn"
TypeMTProxy = "mtproxy"
TypeParser = "parser"
TypeHysteria = "hysteria"
@@ -28,7 +27,6 @@ const (
TypeMieru = "mieru"
TypeAnyTLS = "anytls"
TypeSudoku = "sudoku"
TypeSnell = "snell"
TypeCall = "call"
TypeShadowsocksR = "shadowsocksr"
TypeVLESS = "vless"
@@ -100,8 +98,6 @@ func ProxyDisplayName(proxyType string) string {
return "WARP"
case TypeMASQUE:
return "MASQUE"
case TypeOpenVPN:
return "OpenVPN"
case TypeMTProxy:
return "MTProxy"
case TypeParser:
@@ -132,8 +128,6 @@ func ProxyDisplayName(proxyType string) string {
return "AnyTLS"
case TypeSudoku:
return "Sudoku"
case TypeSnell:
return "Snell"
case TypeCall:
return "Call"
case TypeFallback:
-48
View File
@@ -1,48 +0,0 @@
{
"log": {
"level": "info"
},
"inbounds": [
{
"type": "mixed",
"tag": "mixed-in",
"listen_port": 7897
}
],
"outbounds": [
{
"type": "openvpn",
"tag": "openvpn-out",
"system": false,
"name": "openvpn0",
"servers": [
{
"server": "vpn.example.com",
"server_port": 1194
}
],
"proto": "udp", // udp, tcp
"username": "myuser",
"password": "mypassword",
"tls_crypt": "-----BEGIN OpenVPN Static key V1-----\n...\n-----END OpenVPN Static key V1-----",
// or: "tls_crypt_path": "/path/to/ta.key",
"tls": {
"ca": "-----BEGIN CERTIFICATE-----\n...\n-----END CERTIFICATE-----",
// or: "ca_path": "/path/to/ca.crt",
"cipher_suites": [],
"verify_x509_name": "",
"verify_x509_name_mode": "", // name-prefix, name-suffix, exact (default)
"fragment": false,
"fragment_fallback_delay": "300ms",
"record_fragment": false,
"kernel_tx": false,
"kernel_rx": false
}
// Dial Fields
}
],
"route": {
"final": "openvpn-out",
"auto_detect_interface": true
}
}
-53
View File
@@ -1,53 +0,0 @@
{
"log": {
"level": "info"
},
"inbounds": [
{
"type": "mixed",
"tag": "mixed-in",
"listen_port": 7897
}
],
"outbounds": [
{
"type": "openvpn",
"tag": "openvpn-out",
"system": false,
"name": "openvpn0",
"servers": [
{
"server": "vpn.example.com",
"server_port": 1194
}
],
"proto": "udp", // udp, tcp
"cipher": "AES-256-CBC",
"auth": "SHA1",
"tls_auth": "-----BEGIN OpenVPN Static key V1-----\n...\n-----END OpenVPN Static key V1-----",
// or: "tls_auth_path": "/path/to/ta.key",
"key_direction": 1,
"tls": {
"certificate": "-----BEGIN CERTIFICATE-----\n...\n-----END CERTIFICATE-----",
// or: "certificate_path": "/path/to/client.crt",
"key": "-----BEGIN PRIVATE KEY-----\n...\n-----END PRIVATE KEY-----",
// or: "key_path": "/path/to/client.key",
"ca": "-----BEGIN CERTIFICATE-----\n...\n-----END CERTIFICATE-----",
// or: "ca_path": "/path/to/ca.crt",
"cipher_suites": [],
"verify_x509_name": "",
"verify_x509_name_mode": "", // name-prefix, name-suffix, exact (default)
"fragment": false,
"fragment_fallback_delay": "300ms",
"record_fragment": false,
"kernel_tx": false,
"kernel_rx": false
}
// Dial Fields
}
],
"route": {
"final": "openvpn-out",
"auto_detect_interface": true
}
}
-51
View File
@@ -1,51 +0,0 @@
{
"log": {
"level": "info"
},
"inbounds": [
{
"type": "mixed",
"tag": "mixed-in",
"listen_port": 7897
}
],
"outbounds": [
{
"type": "openvpn",
"tag": "openvpn-out",
"system": false,
"name": "openvpn0",
"servers": [
{
"server": "vpn.example.com",
"server_port": 1194
}
],
"proto": "udp", // udp, tcp
"tls_crypt": "-----BEGIN OpenVPN tls-crypt-v2 client key-----\n...\n-----END OpenVPN tls-crypt-v2 client key-----",
// or: "tls_crypt_path": "/path/to/tls-crypt-v2.key",
"tls_crypt_v2": true,
"tls": {
"certificate": "-----BEGIN CERTIFICATE-----\n...\n-----END CERTIFICATE-----",
// or: "certificate_path": "/path/to/client.crt",
"key": "-----BEGIN PRIVATE KEY-----\n...\n-----END PRIVATE KEY-----",
// or: "key_path": "/path/to/client.key",
"ca": "-----BEGIN CERTIFICATE-----\n...\n-----END CERTIFICATE-----",
// or: "ca_path": "/path/to/ca.crt",
"cipher_suites": [],
"verify_x509_name": "",
"verify_x509_name_mode": "", // name-prefix, name-suffix, exact (default)
"fragment": false,
"fragment_fallback_delay": "300ms",
"record_fragment": false,
"kernel_tx": false,
"kernel_rx": false
}
// Dial Fields
}
],
"route": {
"final": "openvpn-out",
"auto_detect_interface": true
}
}
-54
View File
@@ -1,54 +0,0 @@
{
"log": {
"level": "info"
},
"inbounds": [
{
"type": "mixed",
"tag": "mixed-in",
"listen_port": 7897
}
],
"outbounds": [
{
"type": "openvpn",
"tag": "openvpn-out",
"system": false,
"name": "openvpn0",
"servers": [
{
"server": "vpn.example.com",
"server_port": 1194
}
],
"proto": "udp", // udp, tcp
"cipher": "AES-256-GCM", // AES-128-GCM, AES-192-GCM, AES-256-GCM, AES-128-CBC, AES-192-CBC, AES-256-CBC, CHACHA20-POLY1305
"auth": "SHA256", // SHA1, SHA256, SHA384, SHA512 (ignored for AEAD ciphers)
"tls_crypt": "-----BEGIN OpenVPN Static key V1-----\n...\n-----END OpenVPN Static key V1-----",
// or: "tls_crypt_path": "/path/to/ta.key",
"ping_interval": "10s",
"reconnect_delay": "30s",
"tls": {
"certificate": "-----BEGIN CERTIFICATE-----\n...\n-----END CERTIFICATE-----",
// or: "certificate_path": "/path/to/client.crt",
"key": "-----BEGIN PRIVATE KEY-----\n...\n-----END PRIVATE KEY-----",
// or: "key_path": "/path/to/client.key",
"ca": "-----BEGIN CERTIFICATE-----\n...\n-----END CERTIFICATE-----",
// or: "ca_path": "/path/to/ca.crt",
"cipher_suites": [],
"verify_x509_name": "",
"verify_x509_name_mode": "", // name-prefix, name-suffix, exact (default)
"fragment": false,
"fragment_fallback_delay": "300ms",
"record_fragment": false,
"kernel_tx": false,
"kernel_rx": false
}
// Dial Fields
}
],
"route": {
"final": "openvpn-out",
"auto_detect_interface": true
}
}
-46
View File
@@ -1,46 +0,0 @@
{
"log": {
"level": "error"
},
"dns": {
"servers": [
{
"type": "local",
"tag": "default"
}
]
},
"inbounds": [
{
"type": "mixed",
"tag": "mixed-in",
"listen_port": 7897
}
],
"outbounds": [
{
"type": "direct",
"tag": "direct"
},
{
"type": "snell",
"tag": "snell-out",
"server": "example.com",
"server_port": 8443,
"psk": "your-secret-psk",
"version": 4, // 1 | 2 | 3 | 4 | 5 (v5 falls back to v4)
"reuse": true, // v4 only, reuse pooled connections
"network": ["tcp", "udp"]
// "obfs": {
// "mode": "tls", // tls | http
// "host": "bing.com"
// }
// Dial Fields
}
],
"route": {
"final": "snell-out",
"default_domain_resolver": "default",
"auto_detect_interface": true
}
}
-39
View File
@@ -1,39 +0,0 @@
{
"log": {
"level": "error"
},
"dns": {
"servers": [
{
"type": "local",
"tag": "default"
}
]
},
"inbounds": [
{
"type": "snell",
"tag": "snell-in",
"listen": "::",
"listen_port": 8443,
"psk": "your-secret-psk",
"version": 4, // 4 | 5 (server supports v4/v5 only)
"network": ["tcp", "udp"]
// "obfs": {
// "mode": "tls", // tls | http
// "host": "bing.com"
// }
}
],
"outbounds": [
{
"type": "direct",
"tag": "direct"
}
],
"route": {
"final": "direct",
"default_domain_resolver": "default",
"auto_detect_interface": true
}
}
-12
View File
@@ -1,12 +0,0 @@
//go:build with_openvpn
package include
import (
"github.com/sagernet/sing-box/adapter/outbound"
"github.com/sagernet/sing-box/protocol/openvpn"
)
func registerOpenVPNOutbound(registry *outbound.Registry) {
openvpn.RegisterOutbound(registry)
}
-20
View File
@@ -1,20 +0,0 @@
//go:build !with_openvpn
package include
import (
"context"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/adapter/outbound"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
E "github.com/sagernet/sing/common/exceptions"
)
func registerOpenVPNOutbound(registry *outbound.Registry) {
outbound.Register[option.OpenVPNOutboundOptions](registry, C.TypeOpenVPN, func(ctx context.Context, router adapter.Router, logger log.ContextLogger, tag string, options option.OpenVPNOutboundOptions) (adapter.Outbound, error) {
return nil, E.New(`OpenVPN outbound is not included in this build, rebuild with -tags with_openvpn`)
})
}
-3
View File
@@ -91,7 +91,6 @@ func InboundRegistry() *inbound.Registry {
registerStubForRemovedInbounds(registry)
registerMTProxyInbound(registry)
registerSudokuInbound(registry)
registerSnellInbound(registry)
registerCallInbound(registry)
return registry
@@ -121,7 +120,6 @@ func OutboundRegistry() *outbound.Registry {
mieru.RegisterOutbound(registry)
anytls.RegisterOutbound(registry)
registerMASQUEOutbound(registry)
registerOpenVPNOutbound(registry)
bond.RegisterOutbound(registry)
failover.RegisterOutbound(registry)
@@ -137,7 +135,6 @@ func OutboundRegistry() *outbound.Registry {
registerQUICOutbounds(registry)
registerStubForRemovedOutbounds(registry)
registerSudokuOutbound(registry)
registerSnellOutbound(registry)
registerCallOutbound(registry)
return registry
-17
View File
@@ -1,17 +0,0 @@
//go:build with_snell
package include
import (
"github.com/sagernet/sing-box/adapter/inbound"
"github.com/sagernet/sing-box/adapter/outbound"
"github.com/sagernet/sing-box/protocol/snell"
)
func registerSnellInbound(registry *inbound.Registry) {
snell.RegisterInbound(registry)
}
func registerSnellOutbound(registry *outbound.Registry) {
snell.RegisterOutbound(registry)
}
-27
View File
@@ -1,27 +0,0 @@
//go:build !with_snell
package include
import (
"context"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/adapter/inbound"
"github.com/sagernet/sing-box/adapter/outbound"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
E "github.com/sagernet/sing/common/exceptions"
)
func registerSnellInbound(registry *inbound.Registry) {
inbound.Register[option.SnellInboundOptions](registry, C.TypeSnell, func(ctx context.Context, router adapter.Router, logger log.ContextLogger, tag string, options option.SnellInboundOptions) (adapter.Inbound, error) {
return nil, E.New(`Snell is not included in this build, rebuild with -tags with_snell`)
})
}
func registerSnellOutbound(registry *outbound.Registry) {
outbound.Register[option.SnellOutboundOptions](registry, C.TypeSnell, func(ctx context.Context, router adapter.Router, logger log.ContextLogger, tag string, options option.SnellOutboundOptions) (adapter.Outbound, error) {
return nil, E.New(`Snell is not included in this build, rebuild with -tags with_snell`)
})
}
-48
View File
@@ -1,48 +0,0 @@
package option
import (
"net/netip"
"github.com/sagernet/sing/common/json/badoption"
)
type OpenVPNOutboundOptions struct {
DialerOptions
System bool `json:"system,omitempty"`
Name string `json:"name,omitempty"`
AllowedIPs badoption.Listable[netip.Prefix] `json:"allowed_ips,omitempty"`
Servers []ServerOptions `json:"servers"`
Proto string `json:"proto,omitempty"`
Cipher string `json:"cipher,omitempty"`
Auth string `json:"auth,omitempty"`
Username string `json:"username,omitempty"`
Password string `json:"password,omitempty"`
TLSCrypt string `json:"tls_crypt,omitempty"`
TLSCryptPath string `json:"tls_crypt_path,omitempty"`
TLSCryptV2 bool `json:"tls_crypt_v2,omitempty"`
TLSAuth string `json:"tls_auth,omitempty"`
TLSAuthPath string `json:"tls_auth_path,omitempty"`
KeyDirection int `json:"key_direction,omitempty"`
ReconnectDelay badoption.Duration `json:"reconnect_delay,omitempty"`
PingInterval badoption.Duration `json:"ping_interval,omitempty"`
PingRestart badoption.Duration `json:"ping_restart,omitempty"`
OpenVPNOutboundTLSOptionsContainer
}
type OpenVPNTLSOptions struct {
Certificate string `json:"certificate,omitempty"`
CertificatePath string `json:"certificate_path,omitempty"`
Key string `json:"key,omitempty"`
KeyPath string `json:"key_path,omitempty"`
CA string `json:"ca,omitempty"`
CAPath string `json:"ca_path,omitempty"`
CipherSuites badoption.Listable[string] `json:"cipher_suites,omitempty"`
VerifyX509Name string `json:"verify_x509_name,omitempty"`
VerifyX509NameMode string `json:"verify_x509_name_mode,omitempty"`
KernelTx bool `json:"kernel_tx,omitempty"`
KernelRx bool `json:"kernel_rx,omitempty"`
}
type OpenVPNOutboundTLSOptionsContainer struct {
TLS *OpenVPNTLSOptions `json:"tls,omitempty"`
}
-24
View File
@@ -1,24 +0,0 @@
package option
type SnellOutboundOptions struct {
DialerOptions
ServerOptions
PSK string `json:"psk"`
Version int `json:"version,omitempty"`
Reuse bool `json:"reuse,omitempty"`
Network NetworkList `json:"network,omitempty"`
Obfs *SnellObfsOptions `json:"obfs,omitempty"`
}
type SnellInboundOptions struct {
ListenOptions
PSK string `json:"psk"`
Version int `json:"version,omitempty"`
Network NetworkList `json:"network,omitempty"`
Obfs *SnellObfsOptions `json:"obfs,omitempty"`
}
type SnellObfsOptions struct {
Mode string `json:"mode,omitempty"`
Host string `json:"host,omitempty"`
}
-167
View File
@@ -1,167 +0,0 @@
//go:build with_openvpn
package openvpn
import (
"context"
"net"
"os"
"time"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/adapter/outbound"
"github.com/sagernet/sing-box/common/dialer"
"github.com/sagernet/sing-box/common/tls"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
ovpn "github.com/sagernet/sing-box/transport/openvpn"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/bufio"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/service"
)
func RegisterOutbound(registry *outbound.Registry) {
outbound.Register[option.OpenVPNOutboundOptions](registry, C.TypeOpenVPN, NewOutbound)
}
type Outbound struct {
outbound.Adapter
ctx context.Context
dnsRouter adapter.DNSRouter
logger logger.ContextLogger
tunnel *ovpn.Tunnel
}
func NewOutbound(ctx context.Context, router adapter.Router, logger log.ContextLogger, tag string, options option.OpenVPNOutboundOptions) (adapter.Outbound, error) {
tlsConfig, err := tls.NewOpenVPNClient(ctx, logger, common.PtrValueOrDefault(options.TLS))
if err != nil {
return nil, err
}
var tlsKey []byte
keyDirection := -1
if options.TLSAuth != "" || options.TLSAuthPath != "" {
tlsAuth := options.TLSAuth
if tlsAuth == "" {
data, err := os.ReadFile(options.TLSAuthPath)
if err != nil {
return nil, E.Cause(err, "read tls_auth_path")
}
tlsAuth = string(data)
}
tlsKey = []byte(tlsAuth)
keyDirection = options.KeyDirection
} else {
tlsCrypt := options.TLSCrypt
if tlsCrypt == "" && options.TLSCryptPath != "" {
data, err := os.ReadFile(options.TLSCryptPath)
if err != nil {
return nil, E.Cause(err, "read tls_crypt_path")
}
tlsCrypt = string(data)
}
tlsKey = []byte(tlsCrypt)
}
clientConfig := &ovpn.ClientConfig{
Proto: options.Proto,
Cipher: options.Cipher,
Auth: options.Auth,
Username: options.Username,
Password: options.Password,
KeyDirection: keyDirection,
TLSCrypt: tlsKey,
TLSCryptV2: options.TLSCryptV2,
}
if err := clientConfig.Prepare(); err != nil {
return nil, E.Cause(err, "invalid openvpn config")
}
outboundDialer, err := dialer.New(ctx, options.DialerOptions, true)
if err != nil {
return nil, err
}
o := &Outbound{
Adapter: outbound.NewAdapterWithDialerOptions(C.TypeOpenVPN, tag, []string{N.NetworkTCP, N.NetworkUDP}, options.DialerOptions),
ctx: ctx,
dnsRouter: service.FromContext[adapter.DNSRouter](ctx),
logger: logger,
}
tunnel, err := ovpn.NewTunnel(ctx, logger, ovpn.TunnelOptions{
System: options.System,
Name: options.Name,
CreateDialer: func(interfaceName string) N.Dialer {
return common.Must1(dialer.NewDefault(ctx, option.DialerOptions{
BindInterface: interfaceName,
}))
},
Dialer: outboundDialer,
Servers: options.Servers,
TLSConfig: tlsConfig,
Config: clientConfig,
AllowedAddress: options.AllowedIPs,
ReconnectDelay: time.Duration(options.ReconnectDelay),
PingInterval: time.Duration(options.PingInterval),
PingRestart: time.Duration(options.PingRestart),
})
if err != nil {
return nil, err
}
o.tunnel = tunnel
return o, nil
}
func (o *Outbound) Start(stage adapter.StartStage) error {
if stage != adapter.StartStatePostStart {
return nil
}
return o.tunnel.Start()
}
func (o *Outbound) Close() error {
if o.tunnel != nil {
return o.tunnel.Close()
}
return nil
}
func (o *Outbound) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
switch network {
case N.NetworkTCP:
o.logger.InfoContext(ctx, "outbound connection to ", destination)
case N.NetworkUDP:
o.logger.InfoContext(ctx, "outbound packet connection to ", destination)
}
if destination.IsDomain() {
destinationAddresses, err := o.dnsRouter.Lookup(ctx, destination.Fqdn, adapter.DNSQueryOptions{})
if err != nil {
return nil, err
}
return N.DialSerial(ctx, o.tunnel, network, destination, destinationAddresses)
}
if !destination.Addr.IsValid() {
return nil, E.New("invalid destination: ", destination)
}
return o.tunnel.DialContext(ctx, network, destination)
}
func (o *Outbound) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
o.logger.InfoContext(ctx, "outbound packet connection to ", destination)
if destination.IsDomain() {
destinationAddresses, err := o.dnsRouter.Lookup(ctx, destination.Fqdn, adapter.DNSQueryOptions{})
if err != nil {
return nil, err
}
packetConn, destinationAddress, err := N.ListenSerial(ctx, o.tunnel, destination, destinationAddresses)
if err != nil {
return nil, err
}
if destinationAddress.IsValid() && destination != M.SocksaddrFrom(destinationAddress, destination.Port) {
return bufio.NewNATPacketConn(bufio.NewPacketConn(packetConn), M.SocksaddrFrom(destinationAddress, destination.Port), destination), nil
}
return packetConn, nil
}
return o.tunnel.ListenPacket(ctx, destination)
}
-130
View File
@@ -1,130 +0,0 @@
package snell
import (
"context"
"net"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/adapter/inbound"
"github.com/sagernet/sing-box/common/listener"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing-box/transport/snell"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/bufio"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
func RegisterInbound(registry *inbound.Registry) {
inbound.Register[option.SnellInboundOptions](registry, C.TypeSnell, NewInbound)
}
type Inbound struct {
inbound.Adapter
router adapter.ConnectionRouterEx
logger logger.ContextLogger
listener *listener.Listener
service *snell.Service
}
func NewInbound(ctx context.Context, router adapter.Router, logger log.ContextLogger, tag string, options option.SnellInboundOptions) (adapter.Inbound, error) {
if options.PSK == "" {
return nil, E.New("snell requires psk")
}
udpEnabled := common.Contains(options.Network.Build(), N.NetworkUDP)
obfsMode := ""
if options.Obfs != nil {
obfsMode = options.Obfs.Mode
}
in := &Inbound{
Adapter: inbound.NewAdapter(C.TypeSnell, tag),
router: router,
logger: logger,
}
service, err := snell.NewService(snell.ServiceOptions{
PSK: []byte(options.PSK),
Version: options.Version,
ObfsMode: obfsMode,
UDP: udpEnabled,
Logger: logger,
Handler: (*inboundHandler)(in),
})
if err != nil {
return nil, err
}
in.service = service
in.listener = listener.New(listener.Options{
Context: ctx,
Logger: logger,
Network: []string{N.NetworkTCP},
Listen: options.ListenOptions,
ConnectionHandler: in,
})
return in, nil
}
func (h *Inbound) Start(stage adapter.StartStage) error {
if stage != adapter.StartStateStart {
return nil
}
return h.listener.Start()
}
func (h *Inbound) Close() error {
return h.listener.Close()
}
func (h *Inbound) NewConnectionEx(ctx context.Context, conn net.Conn, metadata adapter.InboundContext, onClose N.CloseHandlerFunc) {
err := h.service.NewConnection(ctx, conn, metadata.Source)
N.CloseOnHandshakeFailure(conn, onClose, err)
if err != nil {
h.logger.ErrorContext(ctx, E.Cause(err, "process connection from ", metadata.Source))
}
}
var _ adapter.TCPInjectableInbound = (*Inbound)(nil)
type inboundHandler Inbound
func (h *inboundHandler) NewConnection(ctx context.Context, conn net.Conn, source M.Socksaddr, destination M.Socksaddr, clientID string) {
var metadata adapter.InboundContext
metadata.Inbound = h.Tag()
metadata.InboundType = h.Type()
metadata.InboundDetour = h.listener.ListenOptions().Detour
metadata.Source = source
metadata.Destination = destination
if clientID != "" {
metadata.User = clientID
h.logger.InfoContext(ctx, "[", clientID, "] inbound connection to ", destination)
} else {
h.logger.InfoContext(ctx, "inbound connection to ", destination)
}
done := make(chan struct{})
h.router.RouteConnectionEx(ctx, conn, metadata, N.OnceClose(func(error) {
close(done)
}))
<-done
}
func (h *inboundHandler) NewPacketConnection(ctx context.Context, conn net.PacketConn, source M.Socksaddr, clientID string) {
var metadata adapter.InboundContext
metadata.Inbound = h.Tag()
metadata.InboundType = h.Type()
metadata.InboundDetour = h.listener.ListenOptions().Detour
metadata.Source = source
if clientID != "" {
metadata.User = clientID
h.logger.InfoContext(ctx, "[", clientID, "] inbound packet connection")
} else {
h.logger.InfoContext(ctx, "inbound packet connection")
}
done := make(chan struct{})
h.router.RoutePacketConnectionEx(ctx, bufio.NewPacketConn(conn), metadata, N.OnceClose(func(error) {
close(done)
}))
<-done
}
-114
View File
@@ -1,114 +0,0 @@
package snell
import (
"context"
"fmt"
"net"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/adapter/outbound"
"github.com/sagernet/sing-box/common/dialer"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/option"
"github.com/sagernet/sing-box/transport/snell"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/bufio"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
func RegisterOutbound(registry *outbound.Registry) {
outbound.Register[option.SnellOutboundOptions](registry, C.TypeSnell, NewOutbound)
}
type Outbound struct {
outbound.Adapter
logger logger.ContextLogger
client *snell.Client
}
func NewOutbound(ctx context.Context, router adapter.Router, logger log.ContextLogger, tag string, options option.SnellOutboundOptions) (adapter.Outbound, error) {
if options.PSK == "" {
return nil, E.New("snell requires psk")
}
version := options.Version
if version == 0 {
version = snell.DefaultSnellVersion
}
if version == snell.Version5 {
version = snell.Version4
}
udpEnabled := common.Contains(options.Network.Build(), N.NetworkUDP)
switch version {
case snell.Version1, snell.Version2:
if udpEnabled {
return nil, fmt.Errorf("snell version %d does not support UDP", version)
}
case snell.Version3, snell.Version4:
default:
return nil, fmt.Errorf("snell version error: %d", version)
}
reuse := version == snell.Version2 || (version == snell.Version4 && options.Reuse)
obfsMode := ""
obfsHost := "bing.com"
if options.Obfs != nil {
switch options.Obfs.Mode {
case "", "tls", "http":
obfsMode = options.Obfs.Mode
default:
return nil, fmt.Errorf("snell obfs mode error: %s", options.Obfs.Mode)
}
if options.Obfs.Host != "" {
obfsHost = options.Obfs.Host
}
}
outboundDialer, err := dialer.New(ctx, options.DialerOptions, options.ServerIsDomain())
if err != nil {
return nil, err
}
client := snell.NewClient(snell.ClientOptions{
Dialer: outboundDialer,
Server: options.ServerOptions.Build(),
PSK: []byte(options.PSK),
Version: version,
Reuse: reuse,
ObfsMode: obfsMode,
ObfsHost: obfsHost,
})
return &Outbound{
Adapter: outbound.NewAdapterWithDialerOptions(C.TypeSnell, tag, options.Network.Build(), options.DialerOptions),
logger: logger,
client: client,
}, nil
}
func (h *Outbound) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
ctx, metadata := adapter.ExtendContext(ctx)
metadata.Outbound = h.Tag()
metadata.Destination = destination
switch N.NetworkName(network) {
case N.NetworkTCP:
h.logger.InfoContext(ctx, "outbound connection to ", destination)
return h.client.DialContext(ctx, destination)
case N.NetworkUDP:
h.logger.InfoContext(ctx, "outbound packet connection to ", destination)
conn, err := h.client.ListenPacket(ctx, destination)
if err != nil {
return nil, err
}
return bufio.NewBindPacketConn(conn, destination), nil
default:
return nil, E.Extend(N.ErrUnknownNetwork, network)
}
}
func (h *Outbound) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
ctx, metadata := adapter.ExtendContext(ctx)
metadata.Outbound = h.Tag()
metadata.Destination = destination
h.logger.InfoContext(ctx, "outbound packet connection to ", destination)
return h.client.ListenPacket(ctx, destination)
}
+1 -1
View File
@@ -1 +1 @@
with_gvisor,with_quic,with_dhcp,with_wireguard,with_utls,with_acme,with_clash_api,with_tailscale,with_masque,with_mtproxy,with_ccm,with_ocm,with_openvpn,with_trusttunnel,with_call,with_sudoku,with_snell,with_naive_outbound,badlinkname,tfogo_checklinkname0
with_gvisor,with_quic,with_dhcp,with_wireguard,with_utls,with_acme,with_clash_api,with_tailscale,with_masque,with_mtproxy,with_ccm,with_ocm,with_trusttunnel,with_call,with_sudoku,with_naive_outbound,badlinkname,tfogo_checklinkname0
+1 -1
View File
@@ -1 +1 @@
with_gvisor,with_quic,with_dhcp,with_wireguard,with_utls,with_acme,with_clash_api,with_tailscale,with_manager,with_admin_panel,with_masque,with_mtproxy,with_ccm,with_ocm,with_openvpn,with_trusttunnel,with_call,with_sudoku,with_snell,with_naive_outbound,with_musl,badlinkname,tfogo_checklinkname0
with_gvisor,with_quic,with_dhcp,with_wireguard,with_utls,with_acme,with_clash_api,with_tailscale,with_manager,with_admin_panel,with_masque,with_mtproxy,with_ccm,with_ocm,with_trusttunnel,with_call,with_sudoku,with_naive_outbound,with_musl,badlinkname,tfogo_checklinkname0
+1 -1
View File
@@ -1 +1 @@
with_gvisor,with_quic,with_dhcp,with_wireguard,with_utls,with_acme,with_clash_api,with_tailscale,with_manager,with_admin_panel,with_masque,with_mtproxy,with_ccm,with_ocm,with_openvpn,with_trusttunnel,with_call,with_sudoku,with_snell,badlinkname,tfogo_checklinkname0
with_gvisor,with_quic,with_dhcp,with_wireguard,with_utls,with_acme,with_clash_api,with_tailscale,with_manager,with_admin_panel,with_masque,with_mtproxy,with_ccm,with_ocm,with_trusttunnel,with_call,with_sudoku,badlinkname,tfogo_checklinkname0
+1 -1
View File
@@ -1 +1 @@
with_gvisor,with_quic,with_dhcp,with_wireguard,with_utls,with_acme,with_clash_api,with_tailscale,with_masque,with_mtproxy,with_ccm,with_ocm,with_openvpn,with_trusttunnel,with_call,with_sudoku,with_snell,with_naive_outbound,with_purego,badlinkname,tfogo_checklinkname0
with_gvisor,with_quic,with_dhcp,with_wireguard,with_utls,with_acme,with_clash_api,with_tailscale,with_masque,with_mtproxy,with_ccm,with_ocm,with_trusttunnel,with_call,with_sudoku,with_naive_outbound,with_purego,badlinkname,tfogo_checklinkname0
-237
View File
@@ -1,237 +0,0 @@
package openvpn
import (
"crypto/aes"
"crypto/cipher"
"crypto/hmac"
"crypto/md5"
"crypto/rand"
"crypto/sha1"
"crypto/sha256"
"crypto/sha512"
"encoding/binary"
"errors"
"hash"
"golang.org/x/crypto/chacha20poly1305"
)
const (
AESGCMTagSize = 16
AESGCMIVSize = 12
CBCIVSize = aes.BlockSize
)
type DataCipher interface {
Encrypt(header []byte, packetID uint32, payload []byte) ([]byte, error)
Decrypt(packet []byte, headerSize int) (plaintext []byte, packetID uint32, err error)
}
type AEADDataCipher struct {
send cipher.AEAD
recv cipher.AEAD
sendImplicitIV [AESGCMIVSize]byte
recvImplicitIV [AESGCMIVSize]byte
}
func NewAEADCipher(keys *KeyMaterial, cipherName string) (*AEADDataCipher, error) {
var send, recv cipher.AEAD
var err error
if cipherName == CipherCHACHA20POLY {
send, err = chacha20poly1305.New(keys.SendCipherKey)
if err != nil {
return nil, err
}
recv, err = chacha20poly1305.New(keys.RecvCipherKey)
if err != nil {
return nil, err
}
} else {
sendBlock, err := aes.NewCipher(keys.SendCipherKey)
if err != nil {
return nil, err
}
recvBlock, err := aes.NewCipher(keys.RecvCipherKey)
if err != nil {
return nil, err
}
send, err = cipher.NewGCMWithTagSize(sendBlock, AESGCMTagSize)
if err != nil {
return nil, err
}
recv, err = cipher.NewGCMWithTagSize(recvBlock, AESGCMTagSize)
if err != nil {
return nil, err
}
}
if len(keys.SendHMACKey) < AESGCMIVSize-4 || len(keys.RecvHMACKey) < AESGCMIVSize-4 {
return nil, errors.New("openvpn implicit IV keys are too short")
}
g := &AEADDataCipher{send: send, recv: recv}
copy(g.sendImplicitIV[4:], keys.SendHMACKey[:AESGCMIVSize-4])
copy(g.recvImplicitIV[4:], keys.RecvHMACKey[:AESGCMIVSize-4])
return g, nil
}
func (g *AEADDataCipher) Encrypt(header []byte, packetID uint32, payload []byte) ([]byte, error) {
var pidBytes [4]byte
binary.BigEndian.PutUint32(pidBytes[:], packetID)
nonce := g.nonce(packetID, g.sendImplicitIV)
ad := append(header, pidBytes[:]...)
sealed := g.send.Seal(nil, nonce[:], payload, ad)
out := make([]byte, 0, len(header)+4+len(sealed))
out = append(out, header...)
out = append(out, pidBytes[:]...)
out = append(out, sealed[len(sealed)-AESGCMTagSize:]...)
out = append(out, sealed[:len(sealed)-AESGCMTagSize]...)
return out, nil
}
func (g *AEADDataCipher) Decrypt(packet []byte, headerSize int) ([]byte, uint32, error) {
if len(packet) < headerSize+4+AESGCMTagSize+1 {
return nil, 0, errors.New("openvpn gcm data packet too short")
}
header := packet[:headerSize]
pidBytes := packet[headerSize : headerSize+4]
tag := packet[headerSize+4 : headerSize+4+AESGCMTagSize]
ciphertext := packet[headerSize+4+AESGCMTagSize:]
combined := append(ciphertext, tag...)
ad := append(header, pidBytes...)
packetID := binary.BigEndian.Uint32(pidBytes)
nonce := g.nonce(packetID, g.recvImplicitIV)
plain, err := g.recv.Open(nil, nonce[:], combined, ad)
if err != nil {
return nil, 0, err
}
return plain, packetID, nil
}
func (g *AEADDataCipher) nonce(packetID uint32, implicit [AESGCMIVSize]byte) [AESGCMIVSize]byte {
nonce := implicit
binary.BigEndian.PutUint32(nonce[:4], binary.BigEndian.Uint32(nonce[:4])^packetID)
return nonce
}
type CBCDataCipher struct {
sendBlock cipher.Block
recvBlock cipher.Block
sendHMAC []byte
recvHMAC []byte
newHash func() hash.Hash
hmacSize int
}
func NewCBCCipher(keys *KeyMaterial, auth string) (*CBCDataCipher, error) {
sendBlock, err := aes.NewCipher(keys.SendCipherKey)
if err != nil {
return nil, err
}
recvBlock, err := aes.NewCipher(keys.RecvCipherKey)
if err != nil {
return nil, err
}
var newHash func() hash.Hash
var hmacSize int
switch auth {
case AuthMD5:
newHash = md5.New
hmacSize = md5.Size
case AuthSHA256:
newHash = sha256.New
hmacSize = sha256.Size
case AuthSHA384:
newHash = sha512.New384
hmacSize = 48
case AuthSHA512:
newHash = sha512.New
hmacSize = sha512.Size
default:
newHash = sha1.New
hmacSize = sha1.Size
}
return &CBCDataCipher{
sendBlock: sendBlock,
recvBlock: recvBlock,
sendHMAC: cloneBytes(keys.SendHMACKey[:hmacSize]),
recvHMAC: cloneBytes(keys.RecvHMACKey[:hmacSize]),
newHash: newHash,
hmacSize: hmacSize,
}, nil
}
func (c *CBCDataCipher) Encrypt(header []byte, packetID uint32, payload []byte) ([]byte, error) {
var pidBytes [4]byte
binary.BigEndian.PutUint32(pidBytes[:], packetID)
plain := append(pidBytes[:], payload...)
padLen := aes.BlockSize - (len(plain) % aes.BlockSize)
for i := 0; i < padLen; i++ {
plain = append(plain, byte(padLen))
}
iv := make([]byte, CBCIVSize)
if _, err := rand.Read(iv); err != nil {
return nil, err
}
ct := make([]byte, len(plain))
cipher.NewCBCEncrypter(c.sendBlock, iv).CryptBlocks(ct, plain)
mac := hmac.New(c.newHash, c.sendHMAC)
mac.Write(iv)
mac.Write(ct)
tag := mac.Sum(nil)
out := make([]byte, 0, len(header)+c.hmacSize+CBCIVSize+len(ct))
out = append(out, header...)
out = append(out, tag...)
out = append(out, iv...)
out = append(out, ct...)
return out, nil
}
func (c *CBCDataCipher) Decrypt(packet []byte, headerSize int) ([]byte, uint32, error) {
minSize := headerSize + c.hmacSize + CBCIVSize + aes.BlockSize
if len(packet) < minSize {
return nil, 0, errors.New("openvpn cbc data packet too short")
}
tag := packet[headerSize : headerSize+c.hmacSize]
iv := packet[headerSize+c.hmacSize : headerSize+c.hmacSize+CBCIVSize]
ct := packet[headerSize+c.hmacSize+CBCIVSize:]
if len(ct)%aes.BlockSize != 0 {
return nil, 0, errors.New("openvpn cbc ciphertext not block-aligned")
}
mac := hmac.New(c.newHash, c.recvHMAC)
mac.Write(iv)
mac.Write(ct)
if !hmac.Equal(tag, mac.Sum(nil)) {
return nil, 0, errors.New("openvpn cbc hmac verification failed")
}
plain := make([]byte, len(ct))
cipher.NewCBCDecrypter(c.recvBlock, iv).CryptBlocks(plain, ct)
padLen := int(plain[len(plain)-1])
if padLen < 1 || padLen > aes.BlockSize {
return nil, 0, errors.New("openvpn cbc invalid padding")
}
plain = plain[:len(plain)-padLen]
if len(plain) < 4 {
return nil, 0, errors.New("openvpn cbc payload too short")
}
packetID := binary.BigEndian.Uint32(plain[:4])
return plain[4:], packetID, nil
}
func CipherKeyLength(cipher string) int {
switch cipher {
case CipherAES128GCM, CipherAES128CBC:
return 16
case CipherAES192GCM, CipherAES192CBC:
return 24
default:
return 32
}
}
func IsAEAD(cipher string) bool {
switch cipher {
case CipherAES128GCM, CipherAES192GCM, CipherAES256GCM, CipherCHACHA20POLY:
return true
default:
return false
}
}
-312
View File
@@ -1,312 +0,0 @@
package openvpn
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"net"
"strings"
"sync/atomic"
"time"
"github.com/sagernet/sing/common/tls"
)
const (
defaultHandshakeTimeout = 30 * time.Second
controlRetransmitDelay = time.Second
)
type Client struct {
config *ClientConfig
tlsConfig tls.Config
mux *PacketMux
control *ControlChannel
tlsConn tls.Conn
data *DataChannel
push *PushReply
cancel context.CancelFunc
lastReceiveNano atomic.Int64
}
func NewClient(config *ClientConfig, io PacketIO, tlsConfig tls.Config) (*Client, error) {
if config == nil {
return nil, errors.New("nil openvpn client config")
}
if io == nil {
return nil, errors.New("nil openvpn packet io")
}
if tlsConfig == nil {
return nil, errors.New("nil openvpn tls config")
}
var crypt ControlCrypt
var err error
if config.TLSAuthKey != nil {
crypt, err = NewTLSAuth(config.TLSAuthKey, config.KeyDirection, config.Auth)
if err != nil {
return nil, err
}
} else if config.TLSCryptKey != nil {
crypt, err = NewTLSCrypt(config.TLSCryptKey, true)
if err != nil {
return nil, err
}
}
local, err := NewSessionID()
if err != nil {
return nil, err
}
runCtx, cancel := context.WithCancel(context.Background())
mux := NewPacketMux(io)
go mux.Run(runCtx)
return &Client{
config: config,
tlsConfig: tlsConfig,
mux: mux,
control: NewControlChannel(mux, crypt, local),
cancel: cancel,
}, nil
}
func (c *Client) Handshake(ctx context.Context) (*PushReply, error) {
if c == nil {
return nil, errors.New("nil openvpn client")
}
if _, ok := ctx.Deadline(); !ok {
var cancel context.CancelFunc
ctx, cancel = context.WithTimeout(ctx, defaultHandshakeTimeout)
defer cancel()
}
if c.config.TLSCryptV2WKc != nil {
if err := c.sendResetV3(ctx); err != nil {
return nil, fmt.Errorf("send hard reset v3: %w", err)
}
} else {
if err := c.control.SendReset(ctx); err != nil {
return nil, fmt.Errorf("send hard reset: %w", err)
}
}
if err := c.waitServerReset(ctx); err != nil {
return nil, err
}
controlConn := NewControlConn(c.control)
tlsConn, err := c.tlsConfig.Client(controlConn)
if err != nil {
return nil, fmt.Errorf("openvpn tls client: %w", err)
}
c.tlsConn = tlsConn
if err := c.tlsConn.HandshakeContext(ctx); err != nil {
return nil, fmt.Errorf("openvpn tls handshake: %w", err)
}
clientRecord, err := NewClientKeyMethod2Record(
InstallScriptOptionsString(c.config.Proto, c.config.Cipher, c.config.Auth),
InstallScriptPeerInfo(c.config.Cipher),
strings.TrimSpace(c.config.Username),
c.config.Password,
)
if err != nil {
return nil, err
}
clientBytes, err := clientRecord.MarshalClient()
if err != nil {
return nil, err
}
if _, err := c.tlsConn.Write(clientBytes); err != nil {
return nil, fmt.Errorf("write key method 2 client record: %w", err)
}
serverRecord, err := c.readServerKeyMethod(ctx)
if err != nil {
return nil, err
}
sources := clientRecord.Sources
sources.Server = serverRecord.Sources.Server
keys, err := DeriveClientKeyMaterial(sources, c.control.LocalSessionID(), c.control.RemoteSessionID(), 32)
if err != nil {
return nil, fmt.Errorf("derive data channel keys: %w", err)
}
if _, err := c.tlsConn.Write([]byte(PushRequest + "\x00")); err != nil {
return nil, fmt.Errorf("write push request: %w", err)
}
push, err := c.readPushReply(ctx)
if err != nil {
return nil, err
}
c.push = push
dataCipher := c.config.Cipher
if push.Cipher != "" {
dataCipher = push.Cipher
}
if dataCipher == "" {
return nil, errors.New("openvpn server did not negotiate a cipher and no cipher configured")
}
keyLen := CipherKeyLength(dataCipher)
keys.SendCipherKey = keys.SendCipherKey[:keyLen]
keys.RecvCipherKey = keys.RecvCipherKey[:keyLen]
var cipher DataCipher
if IsAEAD(dataCipher) {
cipher, err = NewAEADCipher(keys, dataCipher)
} else {
cipher, err = NewCBCCipher(keys, c.config.Auth)
}
if err != nil {
return nil, err
}
c.data = NewDataChannel(cipher, push.PeerID, push.CompLZO)
c.markReceive()
return push, nil
}
func (c *Client) WriteIPPacket(ctx context.Context, packet []byte) error {
if c.data == nil {
return errors.New("openvpn data channel is not ready")
}
encrypted, err := c.data.Encrypt(packet)
if err != nil {
return err
}
return c.mux.WritePacket(ctx, encrypted)
}
func (c *Client) ReadIPPacket(ctx context.Context) ([]byte, error) {
if c.data == nil {
return nil, errors.New("openvpn data channel is not ready")
}
for {
packet, err := c.mux.ReadDataPacket(ctx)
if err != nil {
return nil, err
}
plain, err := c.data.Decrypt(packet)
if err != nil {
continue
}
c.markReceive()
return plain, nil
}
}
func (c *Client) SinceReceive() time.Duration {
return time.Duration(int64(time.Since(clientStart)) - c.lastReceiveNano.Load())
}
func (c *Client) markReceive() {
c.lastReceiveNano.Store(int64(time.Since(clientStart)))
}
var clientStart = time.Now().Add(-time.Hour)
func (c *Client) Close() error {
if c.cancel != nil {
c.cancel()
}
if c.tlsConn != nil {
_ = c.tlsConn.Close()
}
if c.mux != nil {
return c.mux.Close()
}
return nil
}
func (c *Client) waitServerReset(ctx context.Context) error {
retransmits := 0
for {
readCtx := ctx
cancel := func() {}
if c.config.Proto == ProtoUDP {
readCtx, cancel = context.WithTimeout(ctx, controlRetransmitDelay)
}
packet, err := c.control.Read(readCtx)
cancel()
if err != nil {
if c.config.Proto == ProtoUDP && errors.Is(err, context.DeadlineExceeded) && ctx.Err() == nil {
if err := c.control.RetransmitPending(ctx); err != nil {
return fmt.Errorf("retransmit hard reset: %w", err)
}
retransmits++
continue
}
return fmt.Errorf("read hard reset response after %d retransmits: %w", retransmits, err)
}
switch packet.Opcode {
case PControlHardResetServerV2:
return c.control.SendAck(ctx)
case PControlHardResetServerV1:
return fmt.Errorf("openvpn server replied with unsupported key method 1 reset")
}
}
}
func (c *Client) readServerKeyMethod(ctx context.Context) (*KeyMethod2Record, error) {
var buf []byte
tmp := make([]byte, 4096)
for {
n, err := c.tlsConn.Read(tmp)
if err != nil {
return nil, fmt.Errorf("read key method 2 server record: %w", err)
}
buf = append(buf, tmp[:n]...)
record, err := ParseServerKeyMethod2Record(buf)
if err == nil {
return record, nil
}
if !strings.Contains(err.Error(), "truncated") && !errors.Is(err, ioStringEOF) {
return nil, err
}
}
}
func (c *Client) readPushReply(ctx context.Context) (*PushReply, error) {
var buf []byte
tmp := make([]byte, 4096)
for {
n, err := c.tlsConn.Read(tmp)
if err != nil {
if errors.Is(err, io.EOF) && len(buf) > 0 {
break
}
return nil, fmt.Errorf("read push reply: %w", err)
}
buf = append(buf, tmp[:n]...)
if bytes.Contains(buf, []byte("\x00")) || strings.Contains(string(buf), "PUSH_REPLY") {
msg := string(buf)
if idx := strings.IndexByte(msg, 0); idx >= 0 {
msg = msg[:idx]
}
if reply, err := ParsePushReply(msg); err == nil {
return reply, nil
}
}
}
return nil, ctx.Err()
}
func (c *Client) sendResetV3(ctx context.Context) error {
c.control.mu.Lock()
messageID := c.control.sendMessage
c.control.sendMessage++
packet := &ControlPacket{
Opcode: PControlHardResetClientV3,
KeyID: c.control.keyID,
LocalSession: c.control.local,
MessageID: messageID,
}
c.control.pending[messageID] = packet
c.control.mu.Unlock()
encoded, err := c.control.encodeAndWrap(ctx, packet)
if err != nil {
return err
}
encoded = append(encoded, c.config.TLSCryptV2WKc...)
return c.control.io.WritePacket(ctx, encoded)
}
var _ net.Conn = (*ControlConn)(nil)
-176
View File
@@ -1,176 +0,0 @@
package openvpn
import (
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"strings"
)
const (
ProtoUDP = "udp"
ProtoTCP = "tcp"
CipherAES128GCM = "AES-128-GCM"
CipherAES192GCM = "AES-192-GCM"
CipherAES256GCM = "AES-256-GCM"
CipherAES128CBC = "AES-128-CBC"
CipherAES192CBC = "AES-192-CBC"
CipherAES256CBC = "AES-256-CBC"
CipherCHACHA20POLY = "CHACHA20-POLY1305"
AuthMD5 = "MD5"
AuthSHA1 = "SHA1"
AuthSHA256 = "SHA256"
AuthSHA384 = "SHA384"
AuthSHA512 = "SHA512"
)
type ClientConfig struct {
Proto string
Cipher string
Auth string
Username string
Password string
KeyDirection int
TLSCrypt []byte
TLSCryptV2 bool
TLSCryptKey []byte
TLSCryptV2WKc []byte
TLSAuthKey []byte
}
func (c *ClientConfig) Prepare() error {
if c == nil {
return errors.New("nil openvpn client config")
}
c.Proto = normalizeProto(c.Proto)
c.Cipher = strings.ToUpper(strings.TrimSpace(c.Cipher))
if c.Auth == "" {
c.Auth = AuthSHA1
}
c.Auth = strings.ToUpper(strings.TrimSpace(c.Auth))
if c.Proto != ProtoUDP && c.Proto != ProtoTCP {
return fmt.Errorf("unsupported openvpn proto %q: only udp and tcp are supported", c.Proto)
}
if c.Cipher != "" && !isValidCipher(c.Cipher) {
return fmt.Errorf("unsupported openvpn cipher %q", c.Cipher)
}
if !isValidAuth(c.Auth) {
return fmt.Errorf("unsupported openvpn auth %q", c.Auth)
}
if c.TLSCryptV2 {
kc, wkc, err := decodeTLSCryptV2Key(c.TLSCrypt)
if err != nil {
return fmt.Errorf("parse tls-crypt-v2 key: %w", err)
}
c.TLSCryptKey = kc
c.TLSCryptV2WKc = wkc
return nil
}
if len(strings.TrimSpace(string(c.TLSCrypt))) == 0 {
return nil
}
key, err := decodeStaticKey(c.TLSCrypt)
if err != nil {
return fmt.Errorf("parse tls key: %w", err)
}
if c.KeyDirection >= 0 {
c.TLSAuthKey = key
} else {
c.TLSCryptKey = key
}
return nil
}
func normalizeProto(proto string) string {
switch strings.ToLower(strings.TrimSpace(proto)) {
case "", "udp", "udp4":
return ProtoUDP
case "tcp", "tcp-client", "tcp4", "tcp4-client":
return ProtoTCP
default:
return strings.ToLower(strings.TrimSpace(proto))
}
}
func isValidCipher(cipher string) bool {
switch cipher {
case CipherAES128GCM, CipherAES192GCM, CipherAES256GCM,
CipherAES128CBC, CipherAES192CBC, CipherAES256CBC,
CipherCHACHA20POLY:
return true
}
return false
}
func isValidAuth(auth string) bool {
switch auth {
case AuthMD5, AuthSHA1, AuthSHA256, AuthSHA384, AuthSHA512:
return true
}
return false
}
func decodeStaticKey(block []byte) ([]byte, error) {
var hexLines []string
for _, raw := range strings.Split(string(block), "\n") {
line := strings.TrimSpace(raw)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
if strings.HasPrefix(line, "-----BEGIN OpenVPN Static key") || strings.HasPrefix(line, "-----END OpenVPN Static key") {
continue
}
hexLines = append(hexLines, line)
}
encoded := strings.Join(hexLines, "")
key, err := hex.DecodeString(encoded)
if err != nil {
return nil, err
}
if len(key) != 256 {
return nil, fmt.Errorf("invalid static key length %d, expected 256 bytes", len(key))
}
return key, nil
}
func decodeTLSCryptV2Key(block []byte) (kc []byte, wkc []byte, err error) {
data, err := decodePEM(block, "OpenVPN tls-crypt-v2 client key")
if err != nil {
return nil, nil, err
}
if len(data) < 256 {
return nil, nil, fmt.Errorf("tls-crypt-v2 key too short: %d bytes", len(data))
}
return data[:256], data[256:], nil
}
func decodePEM(block []byte, expectedHeader string) ([]byte, error) {
lines := strings.Split(string(block), "\n")
var b64 strings.Builder
inBlock := false
for _, line := range lines {
line = strings.TrimSpace(line)
if strings.Contains(line, "BEGIN") && strings.Contains(line, expectedHeader) {
inBlock = true
continue
}
if strings.Contains(line, "END") && strings.Contains(line, expectedHeader) {
break
}
if inBlock {
b64.WriteString(line)
}
}
if b64.Len() == 0 {
return nil, fmt.Errorf("no %s block found", expectedHeader)
}
return base64Decode(b64.String())
}
func base64Decode(s string) ([]byte, error) {
return base64.StdEncoding.DecodeString(s)
}
-534
View File
@@ -1,534 +0,0 @@
package openvpn
import (
"context"
"errors"
"fmt"
"io"
"net"
"sync"
"time"
)
type PacketIO interface {
ReadPacket(ctx context.Context) ([]byte, error)
WritePacket(ctx context.Context, packet []byte) error
Close() error
LocalAddr() net.Addr
RemoteAddr() net.Addr
}
type ControlChannel struct {
io PacketIO
encode func(*ControlPacket, uint32, uint32) ([]byte, error)
decode func([]byte) (*ControlPacket, uint32, uint32, error)
clock func() time.Time
keyID uint8
local SessionID
remote SessionID
mu sync.Mutex
sendPacketID uint32
sendMessage uint32
recvMessage uint32
ackPending []uint32
pending map[uint32]*ControlPacket
recvPending map[uint32]*ControlPacket
readDeadline time.Time
writeDeadline time.Time
}
func NewControlChannel(io PacketIO, crypt ControlCrypt, local SessionID) *ControlChannel {
ch := &ControlChannel{
io: io,
clock: time.Now,
local: local,
pending: make(map[uint32]*ControlPacket),
recvPending: make(map[uint32]*ControlPacket),
}
if crypt != nil {
ch.encode = func(p *ControlPacket, pid uint32, t uint32) ([]byte, error) {
return EncodeControlPacketCrypt(*p, crypt, pid, t)
}
ch.decode = func(pkt []byte) (*ControlPacket, uint32, uint32, error) {
return DecodeControlPacketCrypt(crypt, pkt)
}
} else {
ch.encode = func(p *ControlPacket, _ uint32, _ uint32) ([]byte, error) {
return EncodeControlPacket(*p)
}
ch.decode = func(pkt []byte) (*ControlPacket, uint32, uint32, error) {
cp, err := DecodeControlPacket(pkt)
return cp, 0, 0, err
}
}
return ch
}
func (c *ControlChannel) LocalSessionID() SessionID {
return c.local
}
func (c *ControlChannel) RemoteSessionID() SessionID {
c.mu.Lock()
defer c.mu.Unlock()
return c.remote
}
func (c *ControlChannel) SetRemoteSessionID(id SessionID) {
c.mu.Lock()
c.remote = id
c.mu.Unlock()
}
func (c *ControlChannel) SendReset(ctx context.Context) error {
_, err := c.Send(ctx, PControlHardResetClientV2, nil)
return err
}
func (c *ControlChannel) Send(ctx context.Context, opcode Opcode, payload []byte) (uint32, error) {
if !opcode.HasMessageID() {
return 0, fmt.Errorf("opcode %s cannot carry a reliable message", opcode)
}
c.mu.Lock()
messageID := c.sendMessage
c.sendMessage++
packet := &ControlPacket{
Opcode: opcode,
KeyID: c.keyID,
LocalSession: c.local,
AckIDs: append([]uint32(nil), c.ackPending...),
AckRemoteSession: c.remote,
MessageID: messageID,
Payload: cloneBytes(payload),
}
c.ackPending = nil
c.pending[messageID] = packet
c.mu.Unlock()
if err := c.writeControlPacket(ctx, packet); err != nil {
return 0, err
}
return messageID, nil
}
func (c *ControlChannel) SendAck(ctx context.Context) error {
c.mu.Lock()
if len(c.ackPending) == 0 {
c.mu.Unlock()
return nil
}
packet := &ControlPacket{
Opcode: PAckV1,
KeyID: c.keyID,
LocalSession: c.local,
AckIDs: append([]uint32(nil), c.ackPending...),
AckRemoteSession: c.remote,
}
c.ackPending = nil
c.mu.Unlock()
return c.writeControlPacket(ctx, packet)
}
func (c *ControlChannel) Read(ctx context.Context) (*ControlPacket, error) {
for {
c.mu.Lock()
if packet, ok := c.recvPending[c.recvMessage]; ok {
delete(c.recvPending, c.recvMessage)
c.recvMessage++
c.mu.Unlock()
return packet, nil
}
c.mu.Unlock()
packet, err := c.readControlPacket(ctx)
if err != nil {
return nil, err
}
var deliver *ControlPacket
sendAck := false
c.mu.Lock()
if c.remote == (SessionID{}) && packet.LocalSession != c.local {
c.remote = packet.LocalSession
}
for _, ackID := range packet.AckIDs {
delete(c.pending, ackID)
}
if packet.Opcode.HasMessageID() {
c.ackPending = appendAck(c.ackPending, packet.MessageID)
}
switch {
case packet.Opcode == PAckV1:
case !packet.Opcode.HasMessageID():
deliver = packet
case packet.MessageID < c.recvMessage:
sendAck = true
case packet.MessageID == c.recvMessage:
deliver = packet
c.recvMessage++
default:
if _, exists := c.recvPending[packet.MessageID]; !exists {
c.recvPending[packet.MessageID] = packet
}
sendAck = true
}
c.mu.Unlock()
if deliver != nil {
return deliver, nil
}
if sendAck {
if err := c.SendAck(ctx); err != nil {
return nil, err
}
}
}
}
func (c *ControlChannel) PendingMessages() int {
c.mu.Lock()
defer c.mu.Unlock()
return len(c.pending)
}
func (c *ControlChannel) RetransmitPending(ctx context.Context) error {
c.mu.Lock()
packets := make([]*ControlPacket, 0, len(c.pending))
for _, packet := range c.pending {
cp := *packet
cp.AckIDs = append([]uint32(nil), c.ackPending...)
cp.AckRemoteSession = c.remote
packets = append(packets, &cp)
}
c.ackPending = nil
c.mu.Unlock()
for _, packet := range packets {
if err := c.writeControlPacket(ctx, packet); err != nil {
return err
}
}
return nil
}
func (c *ControlChannel) writeControlPacket(ctx context.Context, packet *ControlPacket) error {
c.mu.Lock()
c.sendPacketID++
packetID := c.sendPacketID
unixTime := uint32(c.clock().Unix())
deadline := c.writeDeadline
c.mu.Unlock()
if !deadline.IsZero() {
var cancel context.CancelFunc
ctx, cancel = context.WithDeadline(ctx, deadline)
defer cancel()
}
encoded, err := c.encode(packet, packetID, unixTime)
if err != nil {
return err
}
return c.io.WritePacket(ctx, encoded)
}
func (c *ControlChannel) encodeAndWrap(ctx context.Context, packet *ControlPacket) ([]byte, error) {
c.mu.Lock()
c.sendPacketID++
packetID := c.sendPacketID
unixTime := uint32(c.clock().Unix())
c.mu.Unlock()
return c.encode(packet, packetID, unixTime)
}
func (c *ControlChannel) readControlPacket(ctx context.Context) (*ControlPacket, error) {
c.mu.Lock()
deadline := c.readDeadline
c.mu.Unlock()
if !deadline.IsZero() {
var cancel context.CancelFunc
ctx, cancel = context.WithDeadline(ctx, deadline)
defer cancel()
}
raw, err := c.io.ReadPacket(ctx)
if err != nil {
return nil, err
}
packet, _, _, err := c.decode(raw)
return packet, err
}
func (c *ControlChannel) SetDeadline(t time.Time) error {
c.mu.Lock()
c.readDeadline = t
c.writeDeadline = t
c.mu.Unlock()
return nil
}
func (c *ControlChannel) SetReadDeadline(t time.Time) error {
c.mu.Lock()
c.readDeadline = t
c.mu.Unlock()
return nil
}
func (c *ControlChannel) SetWriteDeadline(t time.Time) error {
c.mu.Lock()
c.writeDeadline = t
c.mu.Unlock()
return nil
}
func appendAck(acks []uint32, ack uint32) []uint32 {
for _, existing := range acks {
if existing == ack {
return acks
}
}
return append(acks, ack)
}
type ControlConn struct {
channel *ControlChannel
readBuf []byte
closed bool
mu sync.Mutex
}
func NewControlConn(channel *ControlChannel) *ControlConn {
return &ControlConn{channel: channel}
}
func (c *ControlConn) Read(b []byte) (int, error) {
c.mu.Lock()
if c.closed {
c.mu.Unlock()
return 0, net.ErrClosed
}
if len(c.readBuf) > 0 {
n := copy(b, c.readBuf)
c.readBuf = c.readBuf[n:]
c.mu.Unlock()
return n, nil
}
c.mu.Unlock()
for {
packet, err := c.channel.Read(context.Background())
if err != nil {
return 0, err
}
if packet.Opcode != PControlV1 {
if err := c.channel.SendAck(context.Background()); err != nil {
return 0, err
}
continue
}
if err := c.channel.SendAck(context.Background()); err != nil {
return 0, err
}
if len(packet.Payload) == 0 {
continue
}
n := copy(b, packet.Payload)
if n < len(packet.Payload) {
c.mu.Lock()
c.readBuf = append(c.readBuf, packet.Payload[n:]...)
c.mu.Unlock()
}
return n, nil
}
}
func (c *ControlConn) Write(b []byte) (int, error) {
c.mu.Lock()
if c.closed {
c.mu.Unlock()
return 0, net.ErrClosed
}
c.mu.Unlock()
if _, err := c.channel.Send(context.Background(), PControlV1, b); err != nil {
return 0, err
}
return len(b), nil
}
func (c *ControlConn) Close() error {
c.mu.Lock()
if c.closed {
c.mu.Unlock()
return nil
}
c.closed = true
c.mu.Unlock()
return c.channel.io.Close()
}
func (c *ControlConn) LocalAddr() net.Addr {
return c.channel.io.LocalAddr()
}
func (c *ControlConn) RemoteAddr() net.Addr {
return c.channel.io.RemoteAddr()
}
func (c *ControlConn) SetDeadline(t time.Time) error {
return c.channel.SetDeadline(t)
}
func (c *ControlConn) SetReadDeadline(t time.Time) error {
return c.channel.SetReadDeadline(t)
}
func (c *ControlConn) SetWriteDeadline(t time.Time) error {
return c.channel.SetWriteDeadline(t)
}
type streamPacketIO struct {
conn net.Conn
deadlineMu sync.Mutex
readDeadline time.Time
writeDeadline time.Time
}
type datagramPacketIO struct {
conn net.Conn
deadlineMu sync.Mutex
readDeadline time.Time
writeDeadline time.Time
}
func NewDatagramPacketIO(conn net.Conn) PacketIO {
return &datagramPacketIO{conn: conn}
}
func (d *datagramPacketIO) ReadPacket(ctx context.Context) ([]byte, error) {
if err := setReadDeadlineFromContext(d.conn, ctx, &d.deadlineMu, &d.readDeadline); err != nil {
return nil, err
}
buf := make([]byte, 64*1024)
n, err := d.conn.Read(buf)
if err != nil {
return nil, contextIOError(ctx, err)
}
return buf[:n], nil
}
func (d *datagramPacketIO) WritePacket(ctx context.Context, packet []byte) error {
if err := setWriteDeadlineFromContext(d.conn, ctx, &d.deadlineMu, &d.writeDeadline); err != nil {
return err
}
_, err := d.conn.Write(packet)
return contextIOError(ctx, err)
}
func (d *datagramPacketIO) Close() error {
return d.conn.Close()
}
func (d *datagramPacketIO) LocalAddr() net.Addr {
return d.conn.LocalAddr()
}
func (d *datagramPacketIO) RemoteAddr() net.Addr {
return d.conn.RemoteAddr()
}
func NewTCPPacketIO(conn net.Conn) PacketIO {
return &streamPacketIO{conn: conn}
}
func (s *streamPacketIO) ReadPacket(ctx context.Context) ([]byte, error) {
if err := setReadDeadlineFromContext(s.conn, ctx, &s.deadlineMu, &s.readDeadline); err != nil {
return nil, err
}
var lenBuf [2]byte
if _, err := io.ReadFull(s.conn, lenBuf[:]); err != nil {
return nil, contextIOError(ctx, err)
}
size := int(lenBuf[0])<<8 | int(lenBuf[1])
if size == 0 {
return nil, errors.New("empty openvpn tcp packet")
}
packet := make([]byte, size)
if _, err := io.ReadFull(s.conn, packet); err != nil {
return nil, contextIOError(ctx, err)
}
return packet, nil
}
func (s *streamPacketIO) WritePacket(ctx context.Context, packet []byte) error {
if len(packet) > 0xffff {
return fmt.Errorf("openvpn tcp packet too large: %d", len(packet))
}
if err := setWriteDeadlineFromContext(s.conn, ctx, &s.deadlineMu, &s.writeDeadline); err != nil {
return err
}
frame := make([]byte, 2+len(packet))
frame[0] = byte(len(packet) >> 8)
frame[1] = byte(len(packet))
copy(frame[2:], packet)
_, err := s.conn.Write(frame)
return contextIOError(ctx, err)
}
func (s *streamPacketIO) Close() error {
return s.conn.Close()
}
func (s *streamPacketIO) LocalAddr() net.Addr {
return s.conn.LocalAddr()
}
func (s *streamPacketIO) RemoteAddr() net.Addr {
return s.conn.RemoteAddr()
}
func setReadDeadlineFromContext(conn net.Conn, ctx context.Context, mu *sync.Mutex, current *time.Time) error {
deadline, hasDeadline := ctx.Deadline()
mu.Lock()
defer mu.Unlock()
if current.Equal(deadline) {
return nil
}
if hasDeadline {
if err := conn.SetReadDeadline(deadline); err != nil {
return err
}
} else if err := conn.SetReadDeadline(time.Time{}); err != nil {
return err
}
*current = deadline
return nil
}
func setWriteDeadlineFromContext(conn net.Conn, ctx context.Context, mu *sync.Mutex, current *time.Time) error {
deadline, hasDeadline := ctx.Deadline()
mu.Lock()
defer mu.Unlock()
if current.Equal(deadline) {
return nil
}
if hasDeadline {
if err := conn.SetWriteDeadline(deadline); err != nil {
return err
}
} else if err := conn.SetWriteDeadline(time.Time{}); err != nil {
return err
}
*current = deadline
return nil
}
func contextIOError(ctx context.Context, err error) error {
if err == nil {
return nil
}
var netErr net.Error
if errors.As(err, &netErr) && netErr.Timeout() && ctx.Err() != nil {
return ctx.Err()
}
return err
}
-129
View File
@@ -1,129 +0,0 @@
package openvpn
import (
"errors"
"fmt"
"sync"
)
const (
PeerIDUnset uint32 = 0xffffff
dataChannelReplayWindow = 64
)
type DataChannel struct {
cipher DataCipher
keyID uint8
peerID uint32
compLZO bool
mu sync.Mutex
sendPacketID uint32
recvHighest uint32
recvWindow uint64
recvSeen bool
}
func NewDataChannel(cipher DataCipher, peerID uint32, compLZO bool) *DataChannel {
return &DataChannel{
cipher: cipher,
peerID: peerID,
compLZO: compLZO,
}
}
func (d *DataChannel) Encrypt(packet []byte) ([]byte, error) {
if d.compLZO {
compressed, err := lzo1xCompressSafe(packet)
if err != nil {
return nil, err
}
packet = compressed
}
d.mu.Lock()
d.sendPacketID++
packetID := d.sendPacketID
d.mu.Unlock()
return d.cipher.Encrypt(d.dataHeader(), packetID, packet)
}
func (d *DataChannel) Decrypt(packet []byte) ([]byte, error) {
if len(packet) < 1 {
return nil, errors.New("empty openvpn data packet")
}
opcode, _ := parseOpcodeKeyID(packet[0])
headerSize := 1
if opcode == PDataV2 {
headerSize = 4
}
plain, packetID, err := d.cipher.Decrypt(packet, headerSize)
if err != nil {
return nil, err
}
if err := d.acceptPacketID(packetID); err != nil {
return nil, err
}
if d.compLZO {
return lzo1xDecompressSafe(plain)
}
return plain, nil
}
func (d *DataChannel) dataHeader() []byte {
if d.peerID != PeerIDUnset {
return []byte{
opcodeKeyID(PDataV2, d.keyID),
byte(d.peerID >> 16),
byte(d.peerID >> 8),
byte(d.peerID),
}
}
return []byte{opcodeKeyID(PDataV1, d.keyID)}
}
func (d *DataChannel) acceptPacketID(packetID uint32) error {
d.mu.Lock()
defer d.mu.Unlock()
if !d.recvSeen {
d.recvHighest = packetID
d.recvWindow = 1
d.recvSeen = true
return nil
}
if packetID > d.recvHighest {
shift := packetID - d.recvHighest
if shift >= dataChannelReplayWindow {
d.recvWindow = 1
} else {
d.recvWindow = d.recvWindow<<shift | 1
}
d.recvHighest = packetID
return nil
}
diff := d.recvHighest - packetID
if diff >= dataChannelReplayWindow {
return fmt.Errorf("openvpn replayed data packet id %d", packetID)
}
mask := uint64(1) << diff
if d.recvWindow&mask != 0 {
return fmt.Errorf("openvpn replayed data packet id %d", packetID)
}
d.recvWindow |= mask
return nil
}
func ParsePeerID(options string) uint32 {
for _, field := range splitPushOptions(options) {
if len(field) > len("peer-id ") && field[:len("peer-id ")] == "peer-id " {
var id uint32
if _, err := fmt.Sscanf(field, "peer-id %d", &id); err == nil && id <= PeerIDUnset {
return id
}
}
}
return PeerIDUnset
}
-40
View File
@@ -1,40 +0,0 @@
package openvpn
import (
"context"
"net/netip"
"time"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
wgTun "github.com/sagernet/wireguard-go/tun"
)
type Device interface {
wgTun.Device
N.Dialer
Start() error
}
type DeviceOptions struct {
Context context.Context
Logger logger.ContextLogger
System bool
UDPTimeout time.Duration
CreateDialer func(interfaceName string) N.Dialer
Name string
MTU uint32
Address []netip.Prefix
AllowedAddress []netip.Prefix
}
func NewDevice(options DeviceOptions) (Device, error) {
if !options.System {
return newStackDevice(options)
} else if !tun.WithGVisor {
return newSystemDevice(options)
} else {
return newSystemStackDevice(options)
}
}
-308
View File
@@ -1,308 +0,0 @@
//go:build with_gvisor
package openvpn
import (
"context"
"net"
"net/netip"
"os"
"sync"
"github.com/sagernet/gvisor/pkg/buffer"
"github.com/sagernet/gvisor/pkg/tcpip"
"github.com/sagernet/gvisor/pkg/tcpip/adapters/gonet"
"github.com/sagernet/gvisor/pkg/tcpip/header"
"github.com/sagernet/gvisor/pkg/tcpip/network/ipv4"
"github.com/sagernet/gvisor/pkg/tcpip/network/ipv6"
"github.com/sagernet/gvisor/pkg/tcpip/stack"
"github.com/sagernet/sing-box/log"
"github.com/sagernet/sing-box/transport/wireguard"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing/common/buf"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
wgTun "github.com/sagernet/wireguard-go/tun"
)
type stackDevice struct {
ctx context.Context
logger log.ContextLogger
stack *stack.Stack
mtu uint32
events chan wgTun.Event
outbound chan *stack.PacketBuffer
packetOutbound chan *buf.Buffer
done chan struct{}
closeOnce sync.Once
dispatcher stack.NetworkDispatcher
inet4Address netip.Addr
inet6Address netip.Addr
}
func newStackDevice(options DeviceOptions) (*stackDevice, error) {
tunDevice := &stackDevice{
ctx: options.Context,
logger: options.Logger,
mtu: options.MTU,
events: make(chan wgTun.Event, 1),
outbound: make(chan *stack.PacketBuffer, 256),
packetOutbound: make(chan *buf.Buffer, 256),
done: make(chan struct{}),
}
ipStack, err := tun.NewGVisorStackWithOptions((*wireEndpoint)(tunDevice), stack.NICOptions{}, true)
if err != nil {
return nil, err
}
var (
inet4Address netip.Addr
inet6Address netip.Addr
)
for _, prefix := range options.Address {
addr := tun.AddressFromAddr(prefix.Addr())
protoAddr := tcpip.ProtocolAddress{
AddressWithPrefix: tcpip.AddressWithPrefix{
Address: addr,
PrefixLen: prefix.Bits(),
},
}
if prefix.Addr().Is4() {
inet4Address = prefix.Addr()
tunDevice.inet4Address = inet4Address
protoAddr.Protocol = ipv4.ProtocolNumber
} else {
inet6Address = prefix.Addr()
tunDevice.inet6Address = inet6Address
protoAddr.Protocol = ipv6.ProtocolNumber
}
gErr := ipStack.AddProtocolAddress(tun.DefaultNIC, protoAddr, stack.AddressProperties{})
if gErr != nil {
return nil, E.New("parse local address ", protoAddr.AddressWithPrefix, ": ", gErr.String())
}
}
tunDevice.stack = ipStack
return tunDevice, nil
}
func (w *stackDevice) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
addr := tcpip.FullAddress{
NIC: tun.DefaultNIC,
Port: destination.Port,
Addr: tun.AddressFromAddr(destination.Addr),
}
bind := tcpip.FullAddress{
NIC: tun.DefaultNIC,
}
var networkProtocol tcpip.NetworkProtocolNumber
if destination.IsIPv4() {
if !w.inet4Address.IsValid() {
return nil, E.New("missing IPv4 local address")
}
networkProtocol = header.IPv4ProtocolNumber
bind.Addr = tun.AddressFromAddr(w.inet4Address)
} else {
if !w.inet6Address.IsValid() {
return nil, E.New("missing IPv6 local address")
}
networkProtocol = header.IPv6ProtocolNumber
bind.Addr = tun.AddressFromAddr(w.inet6Address)
}
switch N.NetworkName(network) {
case N.NetworkTCP:
tcpConn, err := wireguard.DialTCPWithBind(ctx, w.stack, bind, addr, networkProtocol)
if err != nil {
return nil, err
}
return tcpConn, nil
case N.NetworkUDP:
udpConn, err := gonet.DialUDP(w.stack, &bind, &addr, networkProtocol)
if err != nil {
return nil, err
}
return udpConn, nil
default:
return nil, E.Extend(N.ErrUnknownNetwork, network)
}
}
func (w *stackDevice) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
bind := tcpip.FullAddress{
NIC: tun.DefaultNIC,
}
var networkProtocol tcpip.NetworkProtocolNumber
if destination.IsIPv4() {
if !w.inet4Address.IsValid() {
return nil, E.New("missing IPv4 local address")
}
networkProtocol = header.IPv4ProtocolNumber
bind.Addr = tun.AddressFromAddr(w.inet4Address)
} else {
if !w.inet6Address.IsValid() {
return nil, E.New("missing IPv6 local address")
}
networkProtocol = header.IPv6ProtocolNumber
bind.Addr = tun.AddressFromAddr(w.inet6Address)
}
udpConn, err := gonet.DialUDP(w.stack, &bind, nil, networkProtocol)
if err != nil {
return nil, err
}
return udpConn, nil
}
func (w *stackDevice) Start() error {
w.events <- wgTun.EventUp
return nil
}
func (w *stackDevice) File() *os.File {
return nil
}
func (w *stackDevice) Read(bufs [][]byte, sizes []int, offset int) (count int, err error) {
select {
case packet, ok := <-w.outbound:
if !ok {
return 0, os.ErrClosed
}
defer packet.DecRef()
var copyN int
/*rangeIterate(packet.Data().AsRange(), func(view *buffer.View) {
copyN += copy(bufs[0][offset+copyN:], view.AsSlice())
})*/
for _, view := range packet.AsSlices() {
copyN += copy(bufs[0][offset+copyN:], view)
}
sizes[0] = copyN
return 1, nil
case packet := <-w.packetOutbound:
defer packet.Release()
sizes[0] = copy(bufs[0][offset:], packet.Bytes())
return 1, nil
case <-w.done:
return 0, os.ErrClosed
}
}
func (w *stackDevice) Write(bufs [][]byte, offset int) (count int, err error) {
for _, b := range bufs {
b = b[offset:]
if len(b) == 0 {
continue
}
var networkProtocol tcpip.NetworkProtocolNumber
switch header.IPVersion(b) {
case header.IPv4Version:
networkProtocol = header.IPv4ProtocolNumber
case header.IPv6Version:
networkProtocol = header.IPv6ProtocolNumber
}
packetBuffer := stack.NewPacketBuffer(stack.PacketBufferOptions{
Payload: buffer.MakeWithData(b),
})
w.dispatcher.DeliverNetworkPacket(networkProtocol, packetBuffer)
packetBuffer.DecRef()
count++
}
return
}
func (w *stackDevice) Flush() error {
return nil
}
func (w *stackDevice) MTU() (int, error) {
return int(w.mtu), nil
}
func (w *stackDevice) Name() (string, error) {
return "sing-box", nil
}
func (w *stackDevice) Events() <-chan wgTun.Event {
return w.events
}
func (w *stackDevice) Close() error {
w.closeOnce.Do(func() {
close(w.done)
close(w.events)
w.stack.Close()
for _, endpoint := range w.stack.CleanupEndpoints() {
endpoint.Abort()
}
w.stack.Wait()
})
return nil
}
func (w *stackDevice) BatchSize() int {
return 1
}
var _ stack.LinkEndpoint = (*wireEndpoint)(nil)
type wireEndpoint stackDevice
func (ep *wireEndpoint) MTU() uint32 {
return ep.mtu
}
func (ep *wireEndpoint) SetMTU(mtu uint32) {
}
func (ep *wireEndpoint) MaxHeaderLength() uint16 {
return 0
}
func (ep *wireEndpoint) LinkAddress() tcpip.LinkAddress {
return ""
}
func (ep *wireEndpoint) SetLinkAddress(addr tcpip.LinkAddress) {
}
func (ep *wireEndpoint) Capabilities() stack.LinkEndpointCapabilities {
return stack.CapabilityRXChecksumOffload
}
func (ep *wireEndpoint) Attach(dispatcher stack.NetworkDispatcher) {
ep.dispatcher = dispatcher
}
func (ep *wireEndpoint) IsAttached() bool {
return ep.dispatcher != nil
}
func (ep *wireEndpoint) Wait() {
}
func (ep *wireEndpoint) ARPHardwareType() header.ARPHardwareType {
return header.ARPHardwareNone
}
func (ep *wireEndpoint) AddHeader(buffer *stack.PacketBuffer) {
}
func (ep *wireEndpoint) ParseHeader(ptr *stack.PacketBuffer) bool {
return true
}
func (ep *wireEndpoint) WritePackets(list stack.PacketBufferList) (int, tcpip.Error) {
for _, packetBuffer := range list.AsSlice() {
packetBuffer.IncRef()
select {
case <-ep.done:
return 0, &tcpip.ErrClosedForSend{}
case ep.outbound <- packetBuffer:
}
}
return list.Len(), nil
}
func (ep *wireEndpoint) Close() {
}
func (ep *wireEndpoint) SetOnCloseAction(f func()) {
}
-13
View File
@@ -1,13 +0,0 @@
//go:build !with_gvisor
package openvpn
import "github.com/sagernet/sing-tun"
func newStackDevice(options DeviceOptions) (Device, error) {
return nil, tun.ErrGVisorNotIncluded
}
func newSystemStackDevice(options DeviceOptions) (Device, error) {
return nil, tun.ErrGVisorNotIncluded
}
-191
View File
@@ -1,191 +0,0 @@
package openvpn
import (
"context"
"errors"
"net"
"net/netip"
"os"
"runtime"
"sync"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-tun"
"github.com/sagernet/sing/common"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/service"
wgTun "github.com/sagernet/wireguard-go/tun"
)
var _ Device = (*systemDevice)(nil)
type systemDevice struct {
options DeviceOptions
dialer N.Dialer
device tun.Tun
batchDevice tun.LinuxTUN
events chan wgTun.Event
closeOnce sync.Once
inet4Address netip.Addr
inet6Address netip.Addr
}
func newSystemDevice(options DeviceOptions) (*systemDevice, error) {
if options.Name == "" {
options.Name = tun.CalculateInterfaceName("openvpn")
}
var inet4Address netip.Addr
var inet6Address netip.Addr
if len(options.Address) > 0 {
if prefix := common.Find(options.Address, func(it netip.Prefix) bool {
return it.Addr().Is4()
}); prefix.IsValid() {
inet4Address = prefix.Addr()
}
}
if len(options.Address) > 0 {
if prefix := common.Find(options.Address, func(it netip.Prefix) bool {
return it.Addr().Is6()
}); prefix.IsValid() {
inet6Address = prefix.Addr()
}
}
return &systemDevice{
options: options,
dialer: options.CreateDialer(options.Name),
events: make(chan wgTun.Event, 1),
inet4Address: inet4Address,
inet6Address: inet6Address,
}, nil
}
func (w *systemDevice) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
return w.dialer.DialContext(ctx, network, destination)
}
func (w *systemDevice) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
return w.dialer.ListenPacket(ctx, destination)
}
func (w *systemDevice) Inet4Address() netip.Addr {
return w.inet4Address
}
func (w *systemDevice) Inet6Address() netip.Addr {
return w.inet6Address
}
func (w *systemDevice) Start() error {
networkManager := service.FromContext[adapter.NetworkManager](w.options.Context)
tunOptions := tun.Options{
Name: w.options.Name,
Inet4Address: common.Filter(w.options.Address, func(it netip.Prefix) bool {
return it.Addr().Is4()
}),
Inet6Address: common.Filter(w.options.Address, func(it netip.Prefix) bool {
return it.Addr().Is6()
}),
MTU: w.options.MTU,
GSO: true,
InterfaceScope: true,
Inet4RouteAddress: common.Filter(w.options.AllowedAddress, func(it netip.Prefix) bool {
return it.Addr().Is4()
}),
Inet6RouteAddress: common.Filter(w.options.AllowedAddress, func(it netip.Prefix) bool {
return it.Addr().Is6()
}),
InterfaceMonitor: networkManager.InterfaceMonitor(),
InterfaceFinder: networkManager.InterfaceFinder(),
Logger: w.options.Logger,
}
if runtime.GOOS == "darwin" {
tunOptions.AutoRoute = true
}
tunInterface, err := tun.New(tunOptions)
if err != nil {
return err
}
err = tunInterface.Start()
if err != nil {
tunInterface.Close()
return err
}
w.options.Logger.Notice("started at ", w.options.Name)
w.device = tunInterface
batchTUN, isBatchTUN := tunInterface.(tun.LinuxTUN)
if isBatchTUN && batchTUN.BatchSize() > 1 {
w.batchDevice = batchTUN
}
w.events <- wgTun.EventUp
return nil
}
func (w *systemDevice) File() *os.File {
return nil
}
func (w *systemDevice) Read(bufs [][]byte, sizes []int, offset int) (count int, err error) {
if w.batchDevice != nil {
count, err = w.batchDevice.BatchRead(bufs, offset-tun.PacketOffset, sizes)
} else {
sizes[0], err = w.device.Read(bufs[0][offset-tun.PacketOffset:])
if err == nil {
count = 1
} else if errors.Is(err, tun.ErrTooManySegments) {
err = wgTun.ErrTooManySegments
}
}
return
}
func (w *systemDevice) Write(bufs [][]byte, offset int) (count int, err error) {
if w.batchDevice != nil {
return w.batchDevice.BatchWrite(bufs, offset)
}
for _, packet := range bufs {
if tun.PacketOffset > 0 {
clear(packet[offset-tun.PacketOffset : offset])
tun.PacketFillHeader(packet[offset-tun.PacketOffset:], tun.PacketIPVersion(packet[offset:]))
}
_, err = w.device.Write(packet[offset-tun.PacketOffset:])
if err != nil {
return
}
}
return
}
func (w *systemDevice) Flush() error {
return nil
}
func (w *systemDevice) MTU() (int, error) {
return int(w.options.MTU), nil
}
func (w *systemDevice) Name() (string, error) {
return w.options.Name, nil
}
func (w *systemDevice) Events() <-chan wgTun.Event {
return w.events
}
func (w *systemDevice) Close() error {
var err error
w.closeOnce.Do(func() {
close(w.events)
if w.device != nil {
err = w.device.Close()
}
})
return err
}
func (w *systemDevice) BatchSize() int {
if w.batchDevice != nil {
return w.batchDevice.BatchSize()
}
return 1
}
-200
View File
@@ -1,200 +0,0 @@
//go:build with_gvisor
package openvpn
import (
"net/netip"
"sync"
"github.com/sagernet/gvisor/pkg/buffer"
"github.com/sagernet/gvisor/pkg/tcpip"
"github.com/sagernet/gvisor/pkg/tcpip/header"
"github.com/sagernet/gvisor/pkg/tcpip/network/ipv4"
"github.com/sagernet/gvisor/pkg/tcpip/network/ipv6"
"github.com/sagernet/gvisor/pkg/tcpip/stack"
"github.com/sagernet/sing-tun"
E "github.com/sagernet/sing/common/exceptions"
)
var _ Device = (*systemStackDevice)(nil)
type systemStackDevice struct {
*systemDevice
stack *stack.Stack
endpoint *systemStackEndpoint
writeBufs [][]byte
closeOnce sync.Once
}
func newSystemStackDevice(options DeviceOptions) (*systemStackDevice, error) {
system, err := newSystemDevice(options)
if err != nil {
return nil, err
}
endpoint := &systemStackEndpoint{
mtu: options.MTU,
done: make(chan struct{}),
}
ipStack, err := tun.NewGVisorStackWithOptions(endpoint, stack.NICOptions{}, true)
if err != nil {
return nil, err
}
for _, prefix := range options.Address {
addr := tun.AddressFromAddr(prefix.Addr())
protoAddr := tcpip.ProtocolAddress{
AddressWithPrefix: tcpip.AddressWithPrefix{
Address: addr,
PrefixLen: prefix.Bits(),
},
}
if prefix.Addr().Is4() {
protoAddr.Protocol = ipv4.ProtocolNumber
} else {
protoAddr.Protocol = ipv6.ProtocolNumber
}
gErr := ipStack.AddProtocolAddress(tun.DefaultNIC, protoAddr, stack.AddressProperties{})
if gErr != nil {
return nil, E.New("parse local address ", protoAddr.AddressWithPrefix, ": ", gErr.String())
}
}
sd := &systemStackDevice{
systemDevice: system,
stack: ipStack,
endpoint: endpoint,
}
endpoint.device = sd
return sd, nil
}
func (w *systemStackDevice) Write(bufs [][]byte, offset int) (count int, err error) {
if w.batchDevice != nil {
w.writeBufs = w.writeBufs[:0]
for _, packet := range bufs {
if !w.writeStack(packet[offset:]) {
w.writeBufs = append(w.writeBufs, packet)
}
}
if len(w.writeBufs) > 0 {
return w.batchDevice.BatchWrite(w.writeBufs, offset)
}
} else {
for _, packet := range bufs {
if !w.writeStack(packet[offset:]) {
if tun.PacketOffset > 0 {
clear(packet[offset-tun.PacketOffset : offset])
tun.PacketFillHeader(packet[offset-tun.PacketOffset:], tun.PacketIPVersion(packet[offset:]))
}
_, err = w.device.Write(packet[offset-tun.PacketOffset:])
}
if err != nil {
return
}
}
}
return
}
func (w *systemStackDevice) Close() error {
var err error
w.closeOnce.Do(func() {
close(w.endpoint.done)
w.stack.Close()
for _, endpoint := range w.stack.CleanupEndpoints() {
endpoint.Abort()
}
w.stack.Wait()
err = w.systemDevice.Close()
})
return err
}
func (w *systemStackDevice) writeStack(packet []byte) bool {
var (
networkProtocol tcpip.NetworkProtocolNumber
destination netip.Addr
)
switch header.IPVersion(packet) {
case header.IPv4Version:
networkProtocol = header.IPv4ProtocolNumber
destination = netip.AddrFrom4(header.IPv4(packet).DestinationAddress().As4())
case header.IPv6Version:
networkProtocol = header.IPv6ProtocolNumber
destination = netip.AddrFrom16(header.IPv6(packet).DestinationAddress().As16())
}
for _, prefix := range w.options.Address {
if prefix.Contains(destination) {
return false
}
}
packetBuffer := stack.NewPacketBuffer(stack.PacketBufferOptions{
Payload: buffer.MakeWithData(packet),
})
w.endpoint.dispatcher.DeliverNetworkPacket(networkProtocol, packetBuffer)
packetBuffer.DecRef()
return true
}
type systemStackEndpoint struct {
mtu uint32
done chan struct{}
device Device
dispatcher stack.NetworkDispatcher
}
func (ep *systemStackEndpoint) MTU() uint32 {
return ep.mtu
}
func (ep *systemStackEndpoint) SetMTU(mtu uint32) {
}
func (ep *systemStackEndpoint) MaxHeaderLength() uint16 {
return 0
}
func (ep *systemStackEndpoint) LinkAddress() tcpip.LinkAddress {
return ""
}
func (ep *systemStackEndpoint) SetLinkAddress(addr tcpip.LinkAddress) {
}
func (ep *systemStackEndpoint) Capabilities() stack.LinkEndpointCapabilities {
return stack.CapabilityRXChecksumOffload
}
func (ep *systemStackEndpoint) Attach(dispatcher stack.NetworkDispatcher) {
ep.dispatcher = dispatcher
}
func (ep *systemStackEndpoint) IsAttached() bool {
return ep.dispatcher != nil
}
func (ep *systemStackEndpoint) Wait() {
}
func (ep *systemStackEndpoint) ARPHardwareType() header.ARPHardwareType {
return header.ARPHardwareNone
}
func (ep *systemStackEndpoint) AddHeader(buffer *stack.PacketBuffer) {
}
func (ep *systemStackEndpoint) ParseHeader(ptr *stack.PacketBuffer) bool {
return true
}
func (ep *systemStackEndpoint) WritePackets(list stack.PacketBufferList) (int, tcpip.Error) {
for _, packetBuffer := range list.AsSlice() {
packet := packetBuffer.ToView().AsSlice()
ep.device.Write([][]byte{packet}, 0)
}
return list.Len(), nil
}
func (ep *systemStackEndpoint) Close() {
}
func (ep *systemStackEndpoint) SetOnCloseAction(f func()) {
}
-250
View File
@@ -1,250 +0,0 @@
package openvpn
import (
"crypto/hmac"
"crypto/md5"
"crypto/rand"
"crypto/sha1"
"encoding/binary"
"errors"
"fmt"
"hash"
)
const (
KeyMethod2 = 2
keySourcePreMasterSize = 48
keySourceRandomSize = 32
maxCipherKeyLength = 64
maxHMACKeyLength = 64
keyBlockSize = 2 * (maxCipherKeyLength + maxHMACKeyLength)
keyExpansionID = "OpenVPN"
)
type KeySource struct {
PreMaster [keySourcePreMasterSize]byte
Random1 [keySourceRandomSize]byte
Random2 [keySourceRandomSize]byte
}
type KeySource2 struct {
Client KeySource
Server KeySource
}
type KeyMaterial struct {
SendCipherKey []byte
SendHMACKey []byte
RecvCipherKey []byte
RecvHMACKey []byte
}
type KeyMethod2Record struct {
Sources KeySource2
Options string
Username string
Password string
PeerInfo string
}
func NewClientKeyMethod2Record(options, peerInfo, username, password string) (*KeyMethod2Record, error) {
var record KeyMethod2Record
if _, err := rand.Read(record.Sources.Client.PreMaster[:]); err != nil {
return nil, err
}
if _, err := rand.Read(record.Sources.Client.Random1[:]); err != nil {
return nil, err
}
if _, err := rand.Read(record.Sources.Client.Random2[:]); err != nil {
return nil, err
}
record.Options = options
record.PeerInfo = peerInfo
record.Username = username
record.Password = password
return &record, nil
}
func (r *KeyMethod2Record) MarshalClient() ([]byte, error) {
if r == nil {
return nil, errors.New("nil key method 2 record")
}
out := make([]byte, 0, 4+1+keySourcePreMasterSize+keySourceRandomSize*2+len(r.Options)+16)
out = binary.BigEndian.AppendUint32(out, 0)
out = append(out, KeyMethod2)
out = append(out, r.Sources.Client.PreMaster[:]...)
out = append(out, r.Sources.Client.Random1[:]...)
out = append(out, r.Sources.Client.Random2[:]...)
out = appendOpenVPNString(out, r.Options)
out = appendOpenVPNString(out, r.Username)
out = appendOpenVPNString(out, r.Password)
out = appendOpenVPNString(out, r.PeerInfo)
return out, nil
}
func ParseServerKeyMethod2Record(packet []byte) (*KeyMethod2Record, error) {
if len(packet) < 4+1+keySourceRandomSize*2 {
return nil, errors.New("key method 2 packet too short")
}
if binary.BigEndian.Uint32(packet[:4]) != 0 {
return nil, errors.New("invalid key method 2 prefix")
}
if packet[4]&0x0f != KeyMethod2 {
return nil, fmt.Errorf("unsupported key method %d", packet[4])
}
offset := 5
record := &KeyMethod2Record{}
copy(record.Sources.Server.Random1[:], packet[offset:offset+keySourceRandomSize])
offset += keySourceRandomSize
copy(record.Sources.Server.Random2[:], packet[offset:offset+keySourceRandomSize])
offset += keySourceRandomSize
var err error
record.Options, offset, err = readOpenVPNString(packet, offset)
if err != nil {
return nil, fmt.Errorf("read options: %w", err)
}
record.Username, offset, _ = readOpenVPNString(packet, offset)
record.Password, offset, _ = readOpenVPNString(packet, offset)
record.PeerInfo, _, _ = readOpenVPNString(packet, offset)
return record, nil
}
func DeriveClientKeyMaterial(sources KeySource2, clientSession, serverSession SessionID, cipherKeyLen int) (*KeyMaterial, error) {
if cipherKeyLen != 16 && cipherKeyLen != 24 && cipherKeyLen != 32 {
return nil, fmt.Errorf("unsupported data cipher key length %d", cipherKeyLen)
}
var master [48]byte
if err := openvpnPRF(
sources.Client.PreMaster[:],
keyExpansionID+" master secret",
sources.Client.Random1[:],
sources.Server.Random1[:],
nil,
nil,
master[:],
); err != nil {
return nil, err
}
keyBlock := make([]byte, keyBlockSize)
if err := openvpnPRF(
master[:],
keyExpansionID+" key expansion",
sources.Client.Random2[:],
sources.Server.Random2[:],
clientSession[:],
serverSession[:],
keyBlock,
); err != nil {
return nil, err
}
clientToServer := keyBlock[:maxCipherKeyLength+maxHMACKeyLength]
serverToClient := keyBlock[maxCipherKeyLength+maxHMACKeyLength:]
return &KeyMaterial{
SendCipherKey: cloneBytes(clientToServer[:cipherKeyLen]),
SendHMACKey: cloneBytes(clientToServer[maxCipherKeyLength : maxCipherKeyLength+maxHMACKeyLength]),
RecvCipherKey: cloneBytes(serverToClient[:cipherKeyLen]),
RecvHMACKey: cloneBytes(serverToClient[maxCipherKeyLength : maxCipherKeyLength+maxHMACKeyLength]),
}, nil
}
func InstallScriptOptionsString(proto, cipher, auth string) string {
protoName := "UDPv4"
if proto == ProtoTCP {
protoName = "TCPv4_CLIENT"
}
keysize := "128"
switch cipher {
case CipherAES192GCM, CipherAES192CBC:
keysize = "192"
case CipherAES256GCM, CipherAES256CBC, CipherCHACHA20POLY:
keysize = "256"
}
return fmt.Sprintf("V4,dev-type tun,link-mtu 1550,tun-mtu 1500,proto %s,cipher %s,auth %s,keysize %s,key-method 2,tls-client", protoName, cipher, auth, keysize)
}
func InstallScriptPeerInfo(cipher string) string {
if cipher != "" {
return "IV_VER=sing-box-openvpn\nIV_PROTO=6\nIV_CIPHERS=" + cipher + "\n"
}
return "IV_VER=sing-box-openvpn\nIV_PROTO=6\nIV_CIPHERS=AES-256-GCM:AES-192-GCM:AES-128-GCM:AES-256-CBC:AES-192-CBC:AES-128-CBC:CHACHA20-POLY1305\n"
}
func appendOpenVPNString(out []byte, s string) []byte {
if s == "" {
return binary.BigEndian.AppendUint16(out, 0)
}
if len(s)+1 > 0xffff {
s = s[:0xfffe]
}
out = binary.BigEndian.AppendUint16(out, uint16(len(s)+1))
out = append(out, s...)
out = append(out, 0)
return out
}
func readOpenVPNString(packet []byte, offset int) (string, int, error) {
if offset+2 > len(packet) {
return "", offset, ioStringEOF
}
size := int(binary.BigEndian.Uint16(packet[offset : offset+2]))
offset += 2
if size == 0 {
return "", offset, nil
}
if offset+size > len(packet) {
return "", offset, ioStringEOF
}
raw := packet[offset : offset+size]
offset += size
if raw[len(raw)-1] == 0 {
raw = raw[:len(raw)-1]
}
return string(raw), offset, nil
}
var ioStringEOF = errors.New("openvpn string truncated")
func openvpnPRF(secret []byte, label string, clientSeed, serverSeed, clientSession, serverSession []byte, out []byte) error {
seed := make([]byte, 0, len(label)+len(clientSeed)+len(serverSeed)+len(clientSession)+len(serverSession))
seed = append(seed, label...)
seed = append(seed, clientSeed...)
seed = append(seed, serverSeed...)
seed = append(seed, clientSession...)
seed = append(seed, serverSession...)
split := (len(secret) + 1) / 2
s1 := secret[:split]
s2 := secret[len(secret)-split:]
md5Out := pHash(md5.New, s1, seed, len(out))
sha1Out := pHash(sha1.New, s2, seed, len(out))
for i := range out {
out[i] = md5Out[i] ^ sha1Out[i]
}
return nil
}
func pHash(newHash func() hash.Hash, secret, seed []byte, size int) []byte {
out := make([]byte, 0, size)
a := hmacSum(newHash, secret, seed)
for len(out) < size {
chunkInput := make([]byte, 0, len(a)+len(seed))
chunkInput = append(chunkInput, a...)
chunkInput = append(chunkInput, seed...)
out = append(out, hmacSum(newHash, secret, chunkInput)...)
a = hmacSum(newHash, secret, a)
}
return out[:size]
}
func hmacSum(newHash func() hash.Hash, key, data []byte) []byte {
mac := hmac.New(newHash, key)
_, _ = mac.Write(data)
return mac.Sum(nil)
}
-48
View File
@@ -1,48 +0,0 @@
package openvpn
import (
"bytes"
"errors"
"github.com/rasky/go-lzo"
)
const (
lzoCompressNone = 0xFA
lzoCompressLZO = 0x66
)
var ErrLZODecompress = errors.New("lzo decompression failed")
func lzo1xDecompressSafe(src []byte) ([]byte, error) {
if len(src) == 0 {
return nil, ErrLZODecompress
}
switch src[0] {
case lzoCompressNone:
if len(src) > 1 {
return src[1:], nil
}
return nil, nil
case lzoCompressLZO:
if len(src) > 1 {
r := bytes.NewReader(src[1:])
out, err := lzo.Decompress1X(r, len(src)-1, 0)
if err != nil {
return nil, ErrLZODecompress
}
return out, nil
}
return nil, nil
default:
return nil, ErrLZODecompress
}
}
func lzo1xCompressSafe(src []byte) ([]byte, error) {
lzoPacket := make([]byte, 1+len(src))
lzoPacket[0] = lzoCompressNone
copy(lzoPacket[1:], src)
return lzoPacket, nil
}
-92
View File
@@ -1,92 +0,0 @@
package openvpn
import (
"context"
"net"
"sync"
)
type PacketMux struct {
io PacketIO
control chan []byte
data chan []byte
done chan struct{}
once sync.Once
}
func NewPacketMux(io PacketIO) *PacketMux {
return &PacketMux{
io: io,
control: make(chan []byte, 64),
data: make(chan []byte, 256),
done: make(chan struct{}),
}
}
func (m *PacketMux) Run(ctx context.Context) {
defer m.Close()
for ctx.Err() == nil {
packet, err := m.io.ReadPacket(ctx)
if err != nil {
return
}
if len(packet) == 0 {
continue
}
opcode, _ := parseOpcodeKeyID(packet[0])
ch := m.data
if opcode.IsControl() {
ch = m.control
}
select {
case ch <- packet:
case <-ctx.Done():
return
case <-m.done:
return
}
}
}
func (m *PacketMux) ReadPacket(ctx context.Context) ([]byte, error) {
select {
case packet := <-m.control:
return packet, nil
case <-ctx.Done():
return nil, ctx.Err()
case <-m.done:
return nil, net.ErrClosed
}
}
func (m *PacketMux) ReadDataPacket(ctx context.Context) ([]byte, error) {
select {
case packet := <-m.data:
return packet, nil
case <-ctx.Done():
return nil, ctx.Err()
case <-m.done:
return nil, net.ErrClosed
}
}
func (m *PacketMux) WritePacket(ctx context.Context, packet []byte) error {
return m.io.WritePacket(ctx, packet)
}
func (m *PacketMux) Close() error {
m.once.Do(func() {
close(m.done)
_ = m.io.Close()
})
return nil
}
func (m *PacketMux) LocalAddr() net.Addr {
return m.io.LocalAddr()
}
func (m *PacketMux) RemoteAddr() net.Addr {
return m.io.RemoteAddr()
}
-254
View File
@@ -1,254 +0,0 @@
package openvpn
import (
"crypto/rand"
"encoding/binary"
"errors"
"fmt"
)
type ControlCrypt interface {
Wrap(header []byte, packetID uint32, unixTime uint32, plaintext []byte) ([]byte, error)
Unwrap(packet []byte) (header []byte, packetID uint32, unixTime uint32, plaintext []byte, err error)
}
const (
KeyIDMask = 0x07
OpcodeShift = 3
PControlHardResetClientV1 Opcode = 1
PControlHardResetServerV1 Opcode = 2
PControlSoftResetV1 Opcode = 3
PControlV1 Opcode = 4
PAckV1 Opcode = 5
PDataV1 Opcode = 6
PControlHardResetClientV2 Opcode = 7
PControlHardResetServerV2 Opcode = 8
PDataV2 Opcode = 9
PControlHardResetClientV3 Opcode = 10
PControlWKCV1 Opcode = 11
SessionIDSize = 8
)
type Opcode uint8
func (o Opcode) String() string {
switch o {
case PControlHardResetClientV1:
return "P_CONTROL_HARD_RESET_CLIENT_V1"
case PControlHardResetServerV1:
return "P_CONTROL_HARD_RESET_SERVER_V1"
case PControlSoftResetV1:
return "P_CONTROL_SOFT_RESET_V1"
case PControlV1:
return "P_CONTROL_V1"
case PAckV1:
return "P_ACK_V1"
case PDataV1:
return "P_DATA_V1"
case PControlHardResetClientV2:
return "P_CONTROL_HARD_RESET_CLIENT_V2"
case PControlHardResetServerV2:
return "P_CONTROL_HARD_RESET_SERVER_V2"
case PDataV2:
return "P_DATA_V2"
case PControlHardResetClientV3:
return "P_CONTROL_HARD_RESET_CLIENT_V3"
case PControlWKCV1:
return "P_CONTROL_WKC_V1"
default:
return "P_UNKNOWN"
}
}
func (o Opcode) IsControl() bool {
switch o {
case PControlHardResetClientV1, PControlHardResetServerV1, PControlSoftResetV1, PControlV1,
PAckV1, PControlHardResetClientV2, PControlHardResetServerV2, PControlHardResetClientV3, PControlWKCV1:
return true
default:
return false
}
}
func (o Opcode) HasMessageID() bool {
return o.IsControl() && o != PAckV1
}
type SessionID [SessionIDSize]byte
func NewSessionID() (SessionID, error) {
var id SessionID
_, err := rand.Read(id[:])
return id, err
}
type ControlPacket struct {
Opcode Opcode
KeyID uint8
LocalSession SessionID
AckIDs []uint32
AckRemoteSession SessionID
MessageID uint32
Payload []byte
}
func opcodeKeyID(opcode Opcode, keyID uint8) byte {
return byte(opcode)<<OpcodeShift | (keyID & KeyIDMask)
}
func parseOpcodeKeyID(b byte) (Opcode, uint8) {
return Opcode(b >> OpcodeShift), b & KeyIDMask
}
func EncodeControlPlain(p ControlPacket) ([]byte, error) {
if !p.Opcode.IsControl() {
return nil, fmt.Errorf("opcode %s is not a control opcode", p.Opcode)
}
if len(p.AckIDs) > 255 {
return nil, fmt.Errorf("too many ack ids: %d", len(p.AckIDs))
}
size := 1 + len(p.AckIDs)*4
if len(p.AckIDs) > 0 {
size += SessionIDSize
}
if p.Opcode.HasMessageID() {
size += 4 + len(p.Payload)
}
out := make([]byte, 0, size)
out = append(out, byte(len(p.AckIDs)))
for _, id := range p.AckIDs {
var b [4]byte
binary.BigEndian.PutUint32(b[:], id)
out = append(out, b[:]...)
}
if len(p.AckIDs) > 0 {
out = append(out, p.AckRemoteSession[:]...)
}
if p.Opcode.HasMessageID() {
var b [4]byte
binary.BigEndian.PutUint32(b[:], p.MessageID)
out = append(out, b[:]...)
out = append(out, p.Payload...)
}
return out, nil
}
func EncodeControlPacket(p ControlPacket) ([]byte, error) {
plain, err := EncodeControlPlain(p)
if err != nil {
return nil, err
}
header := make([]byte, TLSCryptHeaderSize)
header[0] = opcodeKeyID(p.Opcode, p.KeyID)
copy(header[1:], p.LocalSession[:])
out := make([]byte, 0, len(header)+len(plain))
out = append(out, header...)
out = append(out, plain...)
return out, nil
}
func EncodeControlPacketCrypt(p ControlPacket, crypt ControlCrypt, packetID uint32, unixTime uint32) ([]byte, error) {
plain, err := EncodeControlPlain(p)
if err != nil {
return nil, err
}
header := make([]byte, TLSCryptHeaderSize)
header[0] = opcodeKeyID(p.Opcode, p.KeyID)
copy(header[1:], p.LocalSession[:])
return crypt.Wrap(header, packetID, unixTime, plain)
}
func DecodeControlPacket(packet []byte) (*ControlPacket, error) {
if len(packet) < TLSCryptHeaderSize {
return nil, errors.New("control packet too short")
}
opcode, keyID := parseOpcodeKeyID(packet[0])
if !opcode.IsControl() {
return nil, fmt.Errorf("opcode %s is not a control opcode", opcode)
}
var local SessionID
copy(local[:], packet[1:])
plain := packet[TLSCryptHeaderSize:]
ackIDs, ackRemote, messageID, payload, err := DecodeControlPlain(opcode, plain)
if err != nil {
return nil, err
}
return &ControlPacket{
Opcode: opcode,
KeyID: keyID,
LocalSession: local,
AckIDs: ackIDs,
AckRemoteSession: ackRemote,
MessageID: messageID,
Payload: payload,
}, nil
}
func DecodeControlPacketCrypt(crypt ControlCrypt, packet []byte) (*ControlPacket, uint32, uint32, error) {
header, packetID, unixTime, plain, err := crypt.Unwrap(packet)
if err != nil {
return nil, 0, 0, err
}
if len(header) != TLSCryptHeaderSize {
return nil, 0, 0, fmt.Errorf("invalid control header length %d", len(header))
}
opcode, keyID := parseOpcodeKeyID(header[0])
if !opcode.IsControl() {
return nil, 0, 0, fmt.Errorf("opcode %s is not a control opcode", opcode)
}
var local SessionID
copy(local[:], header[1:])
ackIDs, ackRemote, messageID, payload, err := DecodeControlPlain(opcode, plain)
if err != nil {
return nil, 0, 0, err
}
return &ControlPacket{
Opcode: opcode,
KeyID: keyID,
LocalSession: local,
AckIDs: ackIDs,
AckRemoteSession: ackRemote,
MessageID: messageID,
Payload: payload,
}, packetID, unixTime, nil
}
func DecodeControlPlain(opcode Opcode, plain []byte) (ackIDs []uint32, ackRemote SessionID, messageID uint32, payload []byte, err error) {
if len(plain) < 1 {
return nil, SessionID{}, 0, nil, errors.New("control payload too short")
}
ackLen := int(plain[0])
offset := 1
if len(plain) < offset+ackLen*4 {
return nil, SessionID{}, 0, nil, errors.New("control ack array truncated")
}
ackIDs = make([]uint32, ackLen)
for i := 0; i < ackLen; i++ {
ackIDs[i] = binary.BigEndian.Uint32(plain[offset : offset+4])
offset += 4
}
if ackLen > 0 {
if len(plain) < offset+SessionIDSize {
return nil, SessionID{}, 0, nil, errors.New("control ack remote session truncated")
}
copy(ackRemote[:], plain[offset:offset+SessionIDSize])
offset += SessionIDSize
}
if opcode.HasMessageID() {
if len(plain) < offset+4 {
return nil, SessionID{}, 0, nil, errors.New("control message id truncated")
}
messageID = binary.BigEndian.Uint32(plain[offset : offset+4])
offset += 4
payload = cloneBytes(plain[offset:])
} else if len(plain) != offset {
return nil, SessionID{}, 0, nil, errors.New("ack packet has trailing payload")
}
return ackIDs, ackRemote, messageID, payload, nil
}
-163
View File
@@ -1,163 +0,0 @@
package openvpn
import (
"fmt"
"net/netip"
"strconv"
"strings"
)
const PushRequest = "PUSH_REQUEST"
type PushReply struct {
Raw string
Prefixes []netip.Prefix
DNS []netip.Addr
PeerID uint32
Cipher string
Ping uint32
PingRestart uint32
MTU uint32
CompLZO bool
Redirect bool
BlockIPv6 bool
}
func ParsePushReply(message string) (*PushReply, error) {
message = strings.TrimRight(message, "\x00")
if !strings.HasPrefix(message, "PUSH_REPLY") {
return nil, fmt.Errorf("unexpected openvpn push message %q", message)
}
reply := &PushReply{
Raw: message,
PeerID: PeerIDUnset,
}
for _, option := range splitPushOptions(message) {
fields := strings.Fields(option)
if len(fields) == 0 {
continue
}
switch fields[0] {
case "ifconfig":
if len(fields) >= 3 {
prefix, err := parseIPv4Ifconfig(fields[1], fields[2])
if err != nil {
return nil, err
}
reply.Prefixes = append(reply.Prefixes, prefix)
}
case "ifconfig-ipv6":
if len(fields) >= 2 {
prefix, err := netip.ParsePrefix(fields[1])
if err != nil {
return nil, fmt.Errorf("parse pushed ipv6 address %q: %w", fields[1], err)
}
reply.Prefixes = append(reply.Prefixes, prefix)
}
case "dhcp-option":
if len(fields) >= 3 && fields[1] == "DNS" {
if addr, err := netip.ParseAddr(fields[2]); err == nil {
reply.DNS = append(reply.DNS, addr)
}
}
case "peer-id":
if len(fields) >= 2 {
id, err := strconv.ParseUint(fields[1], 10, 24)
if err != nil {
return nil, fmt.Errorf("parse pushed peer-id %q: %w", fields[1], err)
}
reply.PeerID = uint32(id)
}
case "redirect-gateway":
reply.Redirect = true
case "block-ipv6":
reply.BlockIPv6 = true
case "cipher":
if len(fields) >= 2 {
reply.Cipher = fields[1]
}
case "ping":
if len(fields) >= 2 {
if v, err := strconv.ParseUint(fields[1], 10, 32); err == nil {
reply.Ping = uint32(v)
}
}
case "ping-restart":
if len(fields) >= 2 {
if v, err := strconv.ParseUint(fields[1], 10, 32); err == nil {
reply.PingRestart = uint32(v)
}
}
case "tun-mtu":
if len(fields) >= 2 {
if v, err := strconv.ParseUint(fields[1], 10, 32); err == nil {
reply.MTU = uint32(v)
}
}
case "comp-lzo":
reply.CompLZO = true
}
}
if len(reply.Prefixes) == 0 {
return nil, fmt.Errorf("openvpn push reply missing ifconfig address")
}
return reply, nil
}
func splitPushOptions(message string) []string {
message = strings.TrimRight(message, "\x00")
parts := strings.Split(message, ",")
if len(parts) > 0 && parts[0] == "PUSH_REPLY" {
parts = parts[1:]
}
out := parts[:0]
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
out = append(out, part)
}
}
return out
}
func parseIPv4Ifconfig(address, maskOrPeer string) (netip.Prefix, error) {
addr, err := netip.ParseAddr(address)
if err != nil {
return netip.Prefix{}, fmt.Errorf("parse pushed ipv4 address %q: %w", address, err)
}
maskAddr, err := netip.ParseAddr(maskOrPeer)
if err != nil {
return netip.Prefix{}, fmt.Errorf("parse pushed ipv4 mask %q: %w", maskOrPeer, err)
}
if !addr.Is4() || !maskAddr.Is4() {
return netip.Prefix{}, fmt.Errorf("openvpn ifconfig requires ipv4 address and mask")
}
if ones, ok := ipv4MaskSize(maskAddr); ok {
return netip.PrefixFrom(addr, ones), nil
}
// Some servers, including SoftEther/VPNGate in net30/p2p mode, push
// "ifconfig <local> <remote>" rather than "ifconfig <local> <netmask>".
// Use a host prefix for that local tunnel address.
return netip.PrefixFrom(addr, 32), nil
}
func ipv4MaskSize(mask netip.Addr) (int, bool) {
maskBytes := mask.As4()
ones := 0
seenZero := false
for _, b := range maskBytes {
for i := 7; i >= 0; i-- {
if b&(1<<i) == 0 {
seenZero = true
continue
}
if seenZero {
return 0, false
}
ones++
}
}
return ones, true
}
-104
View File
@@ -1,104 +0,0 @@
package openvpn
import (
"crypto/hmac"
"crypto/md5"
"crypto/sha1"
"crypto/sha256"
"crypto/sha512"
"encoding/binary"
"errors"
"fmt"
"hash"
)
type TLSAuth struct {
sendHMACKey []byte
recvHMACKey []byte
newHash func() hash.Hash
hmacSize int
}
func NewTLSAuth(staticKey []byte, keyDirection int, auth string) (*TLSAuth, error) {
if len(staticKey) != staticKeySize {
return nil, fmt.Errorf("invalid tls-auth static key length %d, expected %d", len(staticKey), staticKeySize)
}
key0 := staticKey[:keySlotSize]
key1 := staticKey[keySlotSize:]
var sendSlot, recvSlot []byte
if keyDirection == 1 {
sendSlot = key1
recvSlot = key0
} else {
sendSlot = key0
recvSlot = key1
}
var newHash func() hash.Hash
var hmacSize int
switch auth {
case AuthMD5:
newHash = md5.New
hmacSize = md5.Size
case AuthSHA256:
newHash = sha256.New
hmacSize = sha256.Size
case AuthSHA384:
newHash = sha512.New384
hmacSize = 48
case AuthSHA512:
newHash = sha512.New
hmacSize = sha512.Size
default:
newHash = sha1.New
hmacSize = sha1.Size
}
return &TLSAuth{
sendHMACKey: cloneBytes(sendSlot[64 : 64+hmacSize]),
recvHMACKey: cloneBytes(recvSlot[64 : 64+hmacSize]),
newHash: newHash,
hmacSize: hmacSize,
}, nil
}
func (a *TLSAuth) Wrap(header []byte, packetID uint32, unixTime uint32, plaintext []byte) ([]byte, error) {
if len(header) != TLSCryptHeaderSize {
return nil, fmt.Errorf("invalid tls-auth header length %d, expected %d", len(header), TLSCryptHeaderSize)
}
var pid [TLSCryptPIDSize]byte
binary.BigEndian.PutUint32(pid[:4], packetID)
binary.BigEndian.PutUint32(pid[4:], unixTime)
mac := hmac.New(a.newHash, a.sendHMACKey)
mac.Write(pid[:])
mac.Write(header)
mac.Write(plaintext)
tag := mac.Sum(nil)
out := make([]byte, 0, len(header)+a.hmacSize+TLSCryptPIDSize+len(plaintext))
out = append(out, header...)
out = append(out, tag...)
out = append(out, pid[:]...)
out = append(out, plaintext...)
return out, nil
}
func (a *TLSAuth) Unwrap(packet []byte) (header []byte, packetID uint32, unixTime uint32, plaintext []byte, err error) {
minLen := TLSCryptHeaderSize + a.hmacSize + TLSCryptPIDSize
if len(packet) < minLen {
return nil, 0, 0, nil, errors.New("tls-auth packet too short")
}
header = cloneBytes(packet[:TLSCryptHeaderSize])
tag := packet[TLSCryptHeaderSize : TLSCryptHeaderSize+a.hmacSize]
pidStart := TLSCryptHeaderSize + a.hmacSize
pid := packet[pidStart : pidStart+TLSCryptPIDSize]
plaintext = cloneBytes(packet[pidStart+TLSCryptPIDSize:])
mac := hmac.New(a.newHash, a.recvHMACKey)
mac.Write(pid)
mac.Write(header)
mac.Write(plaintext)
tagCheck := mac.Sum(nil)
if !hmac.Equal(tag, tagCheck) {
return nil, 0, 0, nil, errors.New("tls-auth authentication failed")
}
packetID = binary.BigEndian.Uint32(pid[:4])
unixTime = binary.BigEndian.Uint32(pid[4:])
return header, packetID, unixTime, plaintext, nil
}
-128
View File
@@ -1,128 +0,0 @@
package openvpn
import (
"crypto/aes"
"crypto/cipher"
"crypto/hmac"
"crypto/sha256"
"encoding/binary"
"errors"
"fmt"
)
const (
TLSCryptHeaderSize = 1 + 8
TLSCryptPIDSize = 4 + 4
TLSCryptTagSize = sha256.Size
staticKeySize = 256
keySlotSize = 128
cipherKeySize = 32
hmacKeySize = 32
)
type TLSCrypt struct {
encryptCipherKey []byte
encryptHMACKey []byte
decryptCipherKey []byte
decryptHMACKey []byte
}
func NewTLSCrypt(staticKey []byte, client bool) (*TLSCrypt, error) {
if len(staticKey) != staticKeySize {
return nil, fmt.Errorf("invalid tls-crypt static key length %d, expected %d", len(staticKey), staticKeySize)
}
key0 := staticKey[:keySlotSize]
key1 := staticKey[keySlotSize:]
encrypt := key0
decrypt := key1
if client {
encrypt = key1
decrypt = key0
}
return &TLSCrypt{
encryptCipherKey: cloneBytes(encrypt[:cipherKeySize]),
encryptHMACKey: cloneBytes(encrypt[64 : 64+hmacKeySize]),
decryptCipherKey: cloneBytes(decrypt[:cipherKeySize]),
decryptHMACKey: cloneBytes(decrypt[64 : 64+hmacKeySize]),
}, nil
}
func (c *TLSCrypt) Wrap(header []byte, packetID uint32, unixTime uint32, plaintext []byte) ([]byte, error) {
if len(header) != TLSCryptHeaderSize {
return nil, fmt.Errorf("invalid tls-crypt header length %d, expected %d", len(header), TLSCryptHeaderSize)
}
ad := make([]byte, 0, len(header)+TLSCryptPIDSize)
ad = append(ad, header...)
var pid [TLSCryptPIDSize]byte
binary.BigEndian.PutUint32(pid[:4], packetID)
binary.BigEndian.PutUint32(pid[4:], unixTime)
ad = append(ad, pid[:]...)
tag := c.hmac(c.encryptHMACKey, ad, plaintext)
ciphertext, err := aes256ctr(c.encryptCipherKey, tag[:aes.BlockSize], plaintext)
if err != nil {
return nil, err
}
out := make([]byte, 0, len(ad)+len(tag)+len(ciphertext))
out = append(out, ad...)
out = append(out, tag...)
out = append(out, ciphertext...)
return out, nil
}
func (c *TLSCrypt) Unwrap(packet []byte) (header []byte, packetID uint32, unixTime uint32, plaintext []byte, err error) {
if len(packet) < TLSCryptHeaderSize+TLSCryptPIDSize+TLSCryptTagSize {
return nil, 0, 0, nil, errors.New("tls-crypt packet too short")
}
header = cloneBytes(packet[:TLSCryptHeaderSize])
adEnd := TLSCryptHeaderSize + TLSCryptPIDSize
tagEnd := adEnd + TLSCryptTagSize
ad := packet[:adEnd]
tag := packet[adEnd:tagEnd]
ciphertext := packet[tagEnd:]
plaintext, err = aes256ctr(c.decryptCipherKey, tag[:aes.BlockSize], ciphertext)
if err != nil {
return nil, 0, 0, nil, err
}
tagCheck := c.hmac(c.decryptHMACKey, ad, plaintext)
if !hmac.Equal(tag, tagCheck) {
return nil, 0, 0, nil, errors.New("tls-crypt authentication failed")
}
packetID = binary.BigEndian.Uint32(packet[TLSCryptHeaderSize : TLSCryptHeaderSize+4])
unixTime = binary.BigEndian.Uint32(packet[TLSCryptHeaderSize+4 : adEnd])
return header, packetID, unixTime, plaintext, nil
}
func (c *TLSCrypt) hmac(key []byte, parts ...[]byte) []byte {
mac := hmac.New(sha256.New, key)
for _, part := range parts {
_, _ = mac.Write(part)
}
return mac.Sum(nil)
}
func aes256ctr(key, iv, in []byte) ([]byte, error) {
block, err := aes.NewCipher(key)
if err != nil {
return nil, err
}
out := cloneBytes(in)
cipher.NewCTR(block, iv).XORKeyStream(out, out)
return out, nil
}
func cloneBytes(in []byte) []byte {
out := make([]byte, len(in))
copy(out, in)
return out
}
-345
View File
@@ -1,345 +0,0 @@
package openvpn
import (
"bytes"
"context"
"fmt"
"net"
"net/netip"
"os"
"sync"
"time"
"github.com/sagernet/sing-box/option"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/common/tls"
)
type TunnelOptions struct {
System bool
Name string
CreateDialer func(interfaceName string) N.Dialer
Dialer N.Dialer
Servers []option.ServerOptions
TLSConfig tls.Config
Config *ClientConfig
AllowedAddress []netip.Prefix
UDPTimeout time.Duration
ReconnectDelay time.Duration
PingInterval time.Duration
PingRestart time.Duration
}
type Tunnel struct {
ctx context.Context
cancel context.CancelFunc
logger logger.ContextLogger
options TunnelOptions
device Device
client *Client
mtu uint32
serverIndex int
wg sync.WaitGroup
await chan struct{}
mu sync.Mutex
}
func NewTunnel(ctx context.Context, logger logger.ContextLogger, options TunnelOptions) (*Tunnel, error) {
if options.ReconnectDelay == 0 {
options.ReconnectDelay = 5 * time.Second
}
ctx, cancel := context.WithCancel(ctx)
return &Tunnel{
ctx: ctx,
cancel: cancel,
logger: logger,
options: options,
await: make(chan struct{}),
}, nil
}
func (t *Tunnel) Start() error {
go func() {
client, err := t.getClient()
if err != nil {
t.logger.Error("OpenVPN connect: ", err)
close(t.await)
return
}
t.mtu = 1500
if client.push.MTU > 0 {
t.mtu = client.push.MTU
}
deviceOptions := DeviceOptions{
Context: t.ctx,
Logger: t.logger,
System: t.options.System,
UDPTimeout: t.options.UDPTimeout,
CreateDialer: t.options.CreateDialer,
Name: t.options.Name,
MTU: t.mtu,
Address: client.push.Prefixes,
AllowedAddress: t.options.AllowedAddress,
}
device, err := NewDevice(deviceOptions)
if err != nil {
client.Close()
t.logger.Error("create OpenVPN device: ", err)
close(t.await)
return
}
t.device = device
if err := device.Start(); err != nil {
client.Close()
t.logger.Error("start OpenVPN device: ", err)
close(t.await)
return
}
close(t.await)
t.maintainTunnel()
}()
return nil
}
func (t *Tunnel) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
if err := t.isTunnelInitialized(ctx); err != nil {
return nil, err
}
if !destination.Addr.IsValid() {
return nil, E.Cause(os.ErrInvalid, "invalid non-IP destination")
}
return t.device.DialContext(ctx, network, destination)
}
func (t *Tunnel) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
if err := t.isTunnelInitialized(ctx); err != nil {
return nil, err
}
if !destination.Addr.IsValid() {
return nil, E.Cause(os.ErrInvalid, "invalid non-IP destination")
}
return t.device.ListenPacket(ctx, destination)
}
func (t *Tunnel) Close() error {
t.cancel()
t.mu.Lock()
if t.client != nil {
t.client.Close()
t.client = nil
}
if t.device != nil {
t.device.Close()
t.device = nil
}
t.mu.Unlock()
t.wg.Wait()
return nil
}
func (t *Tunnel) isTunnelInitialized(ctx context.Context) error {
select {
case <-t.await:
case <-ctx.Done():
return ctx.Err()
}
if t.device == nil {
return E.New("endpoint not initialized")
}
return nil
}
func (t *Tunnel) maintainTunnel() {
t.wg.Add(2)
go func() {
defer t.wg.Done()
bufs := make([][]byte, 1)
bufs[0] = make([]byte, t.mtu)
sizes := make([]int, 1)
for t.ctx.Err() == nil {
_, err := t.device.Read(bufs, sizes, 0)
if err != nil {
if t.ctx.Err() != nil {
return
}
continue
}
client, err := t.getClient()
if err != nil {
return
}
if err := client.WriteIPPacket(t.ctx, bufs[0][:sizes[0]]); err != nil {
if t.ctx.Err() != nil {
return
}
}
}
}()
go func() {
defer t.wg.Done()
for t.ctx.Err() == nil {
client, err := t.getClient()
if err != nil {
return
}
packet, err := client.ReadIPPacket(t.ctx)
if err != nil {
if t.ctx.Err() != nil {
return
}
if ok := t.closeClient(client); ok {
t.logger.ErrorContext(t.ctx, fmt.Errorf("connection lost: %v", err))
}
continue
}
if bytes.Equal(packet, pingPayload) {
continue
}
if t.ctx.Err() != nil {
return
}
if t.ctx.Err() != nil {
return
}
if _, err := t.device.Write([][]byte{packet}, 0); err != nil {
return
}
}
}()
pingInterval := t.options.PingInterval
if pingInterval == 0 && t.client != nil && t.client.push.Ping > 0 {
pingInterval = time.Duration(t.client.push.Ping) * time.Second
}
if pingInterval > 0 {
go func() {
ticker := time.NewTicker(pingInterval)
defer ticker.Stop()
for {
select {
case <-t.ctx.Done():
return
case <-ticker.C:
client, err := t.getClient()
if err != nil {
return
}
client.WriteIPPacket(t.ctx, pingPayload)
}
}
}()
}
pingRestart := t.options.PingRestart
if pingRestart == 0 && t.client != nil && t.client.push.PingRestart > 0 {
pingRestart = time.Duration(t.client.push.PingRestart) * time.Second
}
if pingRestart > 0 {
t.wg.Add(1)
go func() {
defer t.wg.Done()
ticker := time.NewTicker(pingRestart)
defer ticker.Stop()
for {
select {
case <-t.ctx.Done():
return
case <-ticker.C:
client, err := t.getClient()
if err != nil {
return
}
if client.SinceReceive() >= pingRestart {
if ok := t.closeClient(client); ok {
t.logger.ErrorContext(t.ctx, fmt.Errorf("ping-restart timeout: no packet received for %s", pingRestart))
}
}
}
}
}()
}
<-t.ctx.Done()
}
func (t *Tunnel) getClient() (*Client, error) {
t.mu.Lock()
defer t.mu.Unlock()
if t.ctx.Err() != nil {
return nil, t.ctx.Err()
}
if t.client != nil {
return t.client, nil
}
timer := time.NewTimer(0)
defer timer.Stop()
for {
t.logger.NoticeContext(t.ctx, "connecting to OpenVPN server")
client, err := t.connect()
if err != nil {
t.logger.ErrorContext(t.ctx, fmt.Errorf("connect failed: %v", err))
timer.Reset(t.options.ReconnectDelay)
select {
case <-t.ctx.Done():
return nil, t.ctx.Err()
case <-timer.C:
}
continue
}
t.client = client
t.logger.NoticeContext(t.ctx, "connected to OpenVPN server")
return client, nil
}
}
func (t *Tunnel) closeClient(client *Client) bool {
t.mu.Lock()
defer t.mu.Unlock()
if client == t.client {
t.client.Close()
t.client = nil
return true
}
return false
}
func (t *Tunnel) connect() (*Client, error) {
config := t.options.Config
server := t.options.Servers[t.serverIndex].Build()
t.serverIndex = (t.serverIndex + 1) % len(t.options.Servers)
connectCtx, cancel := context.WithTimeout(t.ctx, t.options.ReconnectDelay)
defer cancel()
var conn net.Conn
var err error
if config.Proto == ProtoTCP {
conn, err = t.options.Dialer.DialContext(connectCtx, N.NetworkTCP, server)
} else {
conn, err = t.options.Dialer.DialContext(connectCtx, N.NetworkUDP, server)
}
if err != nil {
return nil, fmt.Errorf("dial openvpn server: %w", err)
}
var packetIO PacketIO
if config.Proto == ProtoTCP {
packetIO = NewTCPPacketIO(conn)
} else {
packetIO = NewDatagramPacketIO(conn)
}
client, err := NewClient(config, packetIO, t.options.TLSConfig)
if err != nil {
conn.Close()
return nil, err
}
_, err = client.Handshake(connectCtx)
if err != nil {
client.Close()
return nil, fmt.Errorf("openvpn handshake: %w", err)
}
return client, nil
}
var pingPayload = []byte{
0x2a, 0x18, 0x7b, 0xf3, 0x64, 0x1e, 0xb4, 0xcb,
0x07, 0xed, 0x2d, 0x0a, 0x98, 0x1f, 0xc7, 0x48,
}
-144
View File
@@ -1,144 +0,0 @@
package snell
import (
"encoding/binary"
"net"
"strconv"
)
// SOCKS address types as defined in RFC 1928 section 5.
const (
atypIPv4 = 1
atypDomainName = 3
atypIPv6 = 4
)
// socksAddr represents a SOCKS address as defined in RFC 1928 section 5.
type socksAddr []byte
func (a socksAddr) String() string {
var host, port string
switch a[0] {
case atypDomainName:
hostLen := uint16(a[1])
host = string(a[2 : 2+hostLen])
port = strconv.Itoa((int(a[2+hostLen]) << 8) | int(a[2+hostLen+1]))
case atypIPv4:
host = net.IP(a[1 : 1+net.IPv4len]).String()
port = strconv.Itoa((int(a[1+net.IPv4len]) << 8) | int(a[1+net.IPv4len+1]))
case atypIPv6:
host = net.IP(a[1 : 1+net.IPv6len]).String()
port = strconv.Itoa((int(a[1+net.IPv6len]) << 8) | int(a[1+net.IPv6len+1]))
}
return net.JoinHostPort(host, port)
}
// UDPAddr converts a socksAddr to *net.UDPAddr.
func (a socksAddr) UDPAddr() *net.UDPAddr {
if len(a) == 0 {
return nil
}
switch a[0] {
case atypIPv4:
var ip [net.IPv4len]byte
copy(ip[0:], a[1:1+net.IPv4len])
return &net.UDPAddr{IP: net.IP(ip[:]), Port: int(binary.BigEndian.Uint16(a[1+net.IPv4len : 1+net.IPv4len+2]))}
case atypIPv6:
var ip [net.IPv6len]byte
copy(ip[0:], a[1:1+net.IPv6len])
return &net.UDPAddr{IP: net.IP(ip[:]), Port: int(binary.BigEndian.Uint16(a[1+net.IPv6len : 1+net.IPv6len+2]))}
}
return nil
}
// splitSocksAddr slices a SOCKS address from beginning of b. Returns nil if failed.
func splitSocksAddr(b []byte) socksAddr {
addrLen := 1
if len(b) < addrLen {
return nil
}
switch b[0] {
case atypDomainName:
if len(b) < 2 {
return nil
}
addrLen = 1 + 1 + int(b[1]) + 2
case atypIPv4:
addrLen = 1 + net.IPv4len + 2
case atypIPv6:
addrLen = 1 + net.IPv6len + 2
default:
return nil
}
if len(b) < addrLen {
return nil
}
return b[:addrLen]
}
// parseAddr parses the address in string s. Returns nil if failed.
func parseAddr(s string) socksAddr {
var addr socksAddr
host, port, err := net.SplitHostPort(s)
if err != nil {
return nil
}
if ip := net.ParseIP(host); ip != nil {
if ip4 := ip.To4(); ip4 != nil {
addr = make([]byte, 1+net.IPv4len+2)
addr[0] = atypIPv4
copy(addr[1:], ip4)
} else {
addr = make([]byte, 1+net.IPv6len+2)
addr[0] = atypIPv6
copy(addr[1:], ip)
}
} else {
if len(host) > 255 {
return nil
}
addr = make([]byte, 1+1+len(host)+2)
addr[0] = atypDomainName
addr[1] = byte(len(host))
copy(addr[2:], host)
}
portnum, err := strconv.ParseUint(port, 10, 16)
if err != nil {
return nil
}
addr[len(addr)-2], addr[len(addr)-1] = byte(portnum>>8), byte(portnum)
return addr
}
// parseAddrToSocksAddr parses a socks addr from net.Addr.
// This is a fast path of parseAddr(addr.String()).
func parseAddrToSocksAddr(addr net.Addr) socksAddr {
var hostip net.IP
var port int
switch addr := addr.(type) {
case *net.UDPAddr:
hostip = addr.IP
port = addr.Port
case *net.TCPAddr:
hostip = addr.IP
port = addr.Port
case nil:
return nil
}
if hostip == nil {
return parseAddr(addr.String())
}
var parsed socksAddr
if ip4 := hostip.To4(); ip4.DefaultMask() != nil {
parsed = make([]byte, 1+net.IPv4len+2)
parsed[0] = atypIPv4
copy(parsed[1:], ip4)
binary.BigEndian.PutUint16(parsed[1+net.IPv4len:], uint16(port))
} else {
parsed = make([]byte, 1+net.IPv6len+2)
parsed[0] = atypIPv6
copy(parsed[1:], hostip)
binary.BigEndian.PutUint16(parsed[1+net.IPv6len:], uint16(port))
}
return parsed
}
-56
View File
@@ -1,56 +0,0 @@
package snell
import (
"crypto/aes"
"crypto/cipher"
"golang.org/x/crypto/argon2"
"golang.org/x/crypto/chacha20poly1305"
)
// NewAES128GCM returns the AES-128-GCM cipher used by snell v2/v3.
func NewAES128GCM(psk []byte) Cipher {
return &snellCipher{
psk: psk,
keySize: 16,
makeAEAD: aesGCM,
}
}
// NewChacha20Poly1305 returns the ChaCha20-Poly1305 cipher used by snell v1.
func NewChacha20Poly1305(psk []byte) Cipher {
return &snellCipher{
psk: psk,
keySize: 32,
makeAEAD: chacha20poly1305.New,
}
}
type snellCipher struct {
psk []byte
keySize int
makeAEAD func(key []byte) (cipher.AEAD, error)
}
func (sc *snellCipher) KeySize() int { return sc.keySize }
func (sc *snellCipher) SaltSize() int { return 16 }
func (sc *snellCipher) Encrypter(salt []byte) (cipher.AEAD, error) {
return sc.makeAEAD(snellKDF(sc.psk, salt, sc.KeySize()))
}
func (sc *snellCipher) Decrypter(salt []byte) (cipher.AEAD, error) {
return sc.makeAEAD(snellKDF(sc.psk, salt, sc.KeySize()))
}
func snellKDF(psk, salt []byte, keySize int) []byte {
return argon2.IDKey(psk, salt, 3, 8, 1, 32)[:keySize]
}
func aesGCM(key []byte) (cipher.AEAD, error) {
blk, err := aes.NewCipher(key)
if err != nil {
return nil, err
}
return cipher.NewGCM(blk)
}
-120
View File
@@ -1,120 +0,0 @@
package snell
import (
"context"
"net"
"strconv"
"time"
obfs "github.com/sagernet/sing-box/transport/simple-obfs"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
type ClientOptions struct {
Dialer N.Dialer
Server M.Socksaddr
PSK []byte
Version int
Reuse bool
ObfsMode string
ObfsHost string
}
type Client struct {
dialer N.Dialer
server M.Socksaddr
psk []byte
version int
reuse bool
obfsMode string
obfsHost string
pool *Pool
}
func NewClient(options ClientOptions) *Client {
c := &Client{
dialer: options.Dialer,
server: options.Server,
psk: options.PSK,
version: options.Version,
reuse: options.Reuse,
obfsMode: options.ObfsMode,
obfsHost: options.ObfsHost,
}
if c.reuse {
c.pool = NewPool(func(ctx context.Context) (*Snell, error) {
conn, err := c.dialer.DialContext(ctx, N.NetworkTCP, c.server)
if err != nil {
return nil, err
}
return c.streamConn(conn), nil
})
}
return c
}
func (c *Client) streamConn(conn net.Conn) *Snell {
switch c.obfsMode {
case "tls":
conn = obfs.NewTLSObfs(conn, c.obfsHost)
case "http":
conn = obfs.NewHTTPObfs(conn, c.obfsHost, strconv.Itoa(int(c.server.Port)))
}
return StreamConn(conn, c.psk, c.version)
}
func (c *Client) writeHeader(ctx context.Context, conn net.Conn, destination M.Socksaddr, udp bool) (err error) {
if deadline, ok := ctx.Deadline(); ok {
_ = conn.SetWriteDeadline(deadline)
defer conn.SetWriteDeadline(time.Time{})
}
if udp {
err = WriteUDPHeader(conn, c.version)
if err == nil && c.version >= Version4 {
if sc, ok := conn.(*Snell); ok {
err = sc.ReadReply()
}
}
return
}
err = WriteHeaderWithReuse(conn, destination.AddrString(), uint(destination.Port), c.version, c.reuse)
return
}
func (c *Client) DialContext(ctx context.Context, destination M.Socksaddr) (net.Conn, error) {
if c.reuse {
conn, err := c.pool.Get()
if err != nil {
return nil, err
}
if err = c.writeHeader(ctx, conn, destination, false); err != nil {
_ = conn.Close()
return nil, err
}
return conn, nil
}
conn, err := c.dialer.DialContext(ctx, N.NetworkTCP, c.server)
if err != nil {
return nil, err
}
stream := c.streamConn(conn)
if err = c.writeHeader(ctx, stream, destination, false); err != nil {
_ = conn.Close()
return nil, err
}
return stream, nil
}
func (c *Client) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
conn, err := c.dialer.DialContext(ctx, N.NetworkTCP, c.server)
if err != nil {
return nil, err
}
stream := c.streamConn(conn)
if err = c.writeHeader(ctx, stream, destination, true); err != nil {
_ = conn.Close()
return nil, err
}
return PacketConn(stream), nil
}
-153
View File
@@ -1,153 +0,0 @@
package snell
import (
"context"
"io"
"net"
"runtime"
"sync"
"time"
)
// poolEntry holds a pooled item with its insertion time.
// connPool is a small connection pool with age-based eviction.
// milliseconds
// Pool is a pool of reusable snell connections.
type Pool struct {
pool *connPool
}
func (p *Pool) Get() (net.Conn, error) {
return p.GetContext(context.Background())
}
func (p *Pool) GetContext(ctx context.Context) (net.Conn, error) {
elm, err := p.pool.GetContext(ctx)
if err != nil {
return nil, err
}
return &PoolConn{Snell: elm, pool: p}, nil
}
func (p *Pool) Put(conn *Snell) {
if err := HalfClose(conn); err != nil {
_ = conn.Close()
return
}
p.pool.put(conn)
}
// PoolConn wraps a pooled snell connection and returns it to the pool on Close.
type PoolConn struct {
*Snell
pool *Pool
closeWriteOnce sync.Once
closeWriteErr error
closeOnce sync.Once
closeErr error
}
func (pc *PoolConn) Read(b []byte) (int, error) {
n, err := pc.Snell.Read(b)
if err == ErrZeroChunk {
return n, io.EOF
}
return n, err
}
func (pc *PoolConn) Write(b []byte) (int, error) {
return pc.Snell.Write(b)
}
func (pc *PoolConn) CloseWrite() error {
pc.closeWriteOnce.Do(func() {
pc.closeWriteErr = writeZeroChunk(pc.Snell)
})
return pc.closeWriteErr
}
func (pc *PoolConn) Close() error {
pc.closeOnce.Do(func() {
if err := pc.CloseWrite(); err != nil {
pc.closeErr = err
_ = pc.Snell.Close()
return
}
_ = pc.Snell.Conn.SetReadDeadline(time.Time{})
pc.Snell.reply = false
pc.pool.pool.put(pc.Snell)
})
return pc.closeErr
}
// NewPool creates a new snell connection pool using the given factory.
func NewPool(factory func(context.Context) (*Snell, error)) *Pool {
cp := &connPool{
ch: make(chan *poolEntry, 10),
factory: factory,
maxAge: 15000,
evict: func(item *Snell) {
_ = item.Close()
},
}
p := &Pool{pool: cp}
runtime.SetFinalizer(p, recycle)
return p
}
type poolEntry struct {
elm *Snell
time time.Time
}
type connPool struct {
ch chan *poolEntry
factory func(context.Context) (*Snell, error)
evict func(*Snell)
maxAge int64
}
func (p *connPool) GetContext(ctx context.Context) (*Snell, error) {
now := time.Now()
for {
select {
case item := <-p.ch:
if p.maxAge != 0 && now.Sub(item.time).Milliseconds() > p.maxAge {
if p.evict != nil {
p.evict(item.elm)
}
continue
}
return item.elm, nil
default:
return p.factory(ctx)
}
}
}
func (p *connPool) put(item *Snell) {
e := &poolEntry{
elm: item,
time: time.Now(),
}
select {
case p.ch <- e:
return
default:
if p.evict != nil {
p.evict(item)
}
return
}
}
func recycle(p *Pool) {
for item := range p.pool.ch {
if p.pool.evict != nil {
p.pool.evict(item.elm)
}
}
}
-294
View File
@@ -1,294 +0,0 @@
package snell
import (
"bufio"
"context"
"encoding/binary"
"errors"
"fmt"
"io"
"net"
"sync"
"time"
obfs "github.com/sagernet/sing-box/transport/simple-obfs"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
)
type Handler interface {
NewConnection(ctx context.Context, conn net.Conn, source M.Socksaddr, destination M.Socksaddr, clientID string)
NewPacketConnection(ctx context.Context, conn net.PacketConn, source M.Socksaddr, clientID string)
}
type ServiceOptions struct {
PSK []byte
Version int
ObfsMode string
UDP bool
Logger logger.ContextLogger
Handler Handler
}
type Service struct {
psk []byte
version int
obfsMode string
udp bool
logger logger.ContextLogger
handler Handler
}
func NewService(options ServiceOptions) (*Service, error) {
version := options.Version
if version == 0 {
version = Version4
}
if version != Version4 && version != Version5 {
return nil, fmt.Errorf("snell inbound version %d is not supported", version)
}
if len(options.PSK) == 0 {
return nil, errors.New("snell inbound requires psk")
}
switch options.ObfsMode {
case "", "http", "tls":
default:
return nil, fmt.Errorf("snell inbound obfs mode error: %s", options.ObfsMode)
}
return &Service{
psk: options.PSK,
version: version,
obfsMode: options.ObfsMode,
udp: options.UDP,
logger: options.Logger,
handler: options.Handler,
}, nil
}
func (s *Service) NewConnection(ctx context.Context, rawConn net.Conn, source M.Socksaddr) error {
conn := rawConn
switch s.obfsMode {
case "http":
conn = obfs.NewHTTPObfsServer(conn)
case "tls":
conn = obfs.NewTLSObfsServer(conn)
}
stream := ServerStreamConn(conn, s.psk, s.version)
for {
reuse, err := s.handleRequest(ctx, stream, source)
if err != nil || !reuse {
return err
}
}
}
func (s *Service) handleRequest(ctx context.Context, stream *Snell, source M.Socksaddr) (bool, error) {
br := bufio.NewReader(stream)
version, err := br.ReadByte()
if err != nil {
return false, err
}
if version != Version {
return false, fmt.Errorf("snell invalid protocol version: %d", version)
}
command, err := br.ReadByte()
if err != nil {
return false, err
}
if command == CommandPing {
_, _ = stream.Write([]byte{CommandPong})
return false, nil
}
clientID, err := readClientID(br)
if err != nil {
return false, err
}
switch command {
case CommandConnect, CommandConnectV2:
return s.handleTCP(ctx, stream, br, command == CommandConnectV2, clientID, source)
case CommandUDP:
if !s.udp {
return false, errors.New("snell UDP is disabled")
}
return false, s.handleUDP(ctx, stream, clientID, source)
default:
return false, fmt.Errorf("snell unknown command: %d", command)
}
}
func (s *Service) handleTCP(ctx context.Context, stream *Snell, br *bufio.Reader, reuse bool, clientID string, source M.Socksaddr) (bool, error) {
hostLen, err := br.ReadByte()
if err != nil {
return false, err
}
if hostLen == 0 {
return false, errors.New("snell connect host is empty")
}
hostBytes := make([]byte, int(hostLen))
if _, err := io.ReadFull(br, hostBytes); err != nil {
return false, err
}
var portBytes [2]byte
if _, err := io.ReadFull(br, portBytes[:]); err != nil {
return false, err
}
destination := M.ParseSocksaddrHostPort(string(hostBytes), binary.BigEndian.Uint16(portBytes[:]))
conn := &tcpRequestConn{
Conn: stream,
reader: br,
reuse: reuse,
}
s.handler.NewConnection(ctx, conn, source, destination, clientID)
if !reuse {
return false, nil
}
return true, nil
}
func (s *Service) handleUDP(ctx context.Context, stream *Snell, clientID string, source M.Socksaddr) error {
if _, err := stream.Write([]byte{CommandTunnel}); err != nil {
return err
}
pc := &serverPacketConn{
conn: stream,
writeMu: &sync.Mutex{},
}
s.handler.NewPacketConnection(ctx, pc, source, clientID)
return nil
}
const maxPacketLength = 0x3fff
func readClientID(r *bufio.Reader) (string, error) {
length, err := r.ReadByte()
if err != nil {
return "", err
}
if length == 0 {
return "", nil
}
id := make([]byte, int(length))
if _, err := io.ReadFull(r, id); err != nil {
return "", err
}
return string(id), nil
}
func writeCommandError(w io.Writer, code byte, message string) error {
msg := []byte(message)
if len(msg) > 255 {
msg = msg[:255]
}
buf := make([]byte, 0, 3+len(msg))
buf = append(buf, CommandError, code, byte(len(msg)))
buf = append(buf, msg...)
_, err := w.Write(buf)
return err
}
type tcpRequestConn struct {
net.Conn
reader *bufio.Reader
reuse bool
writeMu sync.Mutex
closeOnce sync.Once
replyWritten bool
}
func (c *tcpRequestConn) Read(p []byte) (int, error) {
n, err := c.reader.Read(p)
if errors.Is(err, ErrZeroChunk) {
err = io.EOF
}
return n, err
}
func (c *tcpRequestConn) Write(p []byte) (int, error) {
c.writeMu.Lock()
defer c.writeMu.Unlock()
if !c.replyWritten {
payload := make([]byte, 1+len(p))
payload[0] = CommandTunnel
copy(payload[1:], p)
if _, err := c.Conn.Write(payload); err != nil {
return 0, err
}
c.replyWritten = true
return len(p), nil
}
return c.Conn.Write(p)
}
func (c *tcpRequestConn) CloseWrite() error {
return c.Close()
}
func (c *tcpRequestConn) Close() error {
var err error
c.closeOnce.Do(func() {
c.writeMu.Lock()
defer c.writeMu.Unlock()
if !c.replyWritten {
err = writeCommandError(c.Conn, 0x65, "Remote EOF")
if !c.reuse {
err = errors.Join(err, c.Conn.Close())
}
return
}
if c.reuse {
_, err = c.Conn.Write(nil)
return
}
err = c.Conn.Close()
})
return err
}
type serverPacketConn struct {
conn *Snell
writeMu *sync.Mutex
readBuf []byte
}
func (c *serverPacketConn) ReadFrom(p []byte) (int, net.Addr, error) {
if c.readBuf == nil {
c.readBuf = make([]byte, maxPacketLength)
}
for {
n, err := c.conn.Read(c.readBuf)
if err != nil {
if errors.Is(err, io.EOF) || errors.Is(err, ErrZeroChunk) {
return 0, nil, io.EOF
}
return 0, nil, err
}
request, err := ParseUDPRequest(c.readBuf[:n])
if err != nil {
return 0, nil, err
}
var destination M.Socksaddr
if request.Ip.IsValid() {
destination = M.SocksaddrFrom(request.Ip, request.Port)
} else {
destination = M.ParseSocksaddrHostPort(request.Host, request.Port)
}
length := copy(p, request.Payload)
if destination.IsFqdn() {
return length, destination, nil
}
return length, destination.UDPAddr(), nil
}
}
func (c *serverPacketConn) WriteTo(p []byte, addr net.Addr) (int, error) {
c.writeMu.Lock()
defer c.writeMu.Unlock()
return WritePacketResponse(c.conn, addr, p)
}
func (c *serverPacketConn) Close() error { return c.conn.Close() }
func (c *serverPacketConn) LocalAddr() net.Addr { return c.conn.LocalAddr() }
func (c *serverPacketConn) SetDeadline(t time.Time) error { return c.conn.SetDeadline(t) }
func (c *serverPacketConn) SetReadDeadline(t time.Time) error { return c.conn.SetReadDeadline(t) }
func (c *serverPacketConn) SetWriteDeadline(t time.Time) error { return c.conn.SetWriteDeadline(t) }
-211
View File
@@ -1,211 +0,0 @@
package snell
import (
"crypto/cipher"
"crypto/rand"
"errors"
"io"
"net"
"github.com/sagernet/sing/common/buf"
)
// payloadSizeMask is the maximum size of payload in bytes.
// 16*1024 - 1
// >= 2+aead.Overhead()+payloadSizeMask+aead.Overhead()
// ErrZeroChunk is returned when a zero-length chunk is read, which snell uses
// as an end-of-stream signal.
var ErrZeroChunk = errors.New("zero chunk")
// Cipher is the AEAD cipher abstraction used by the shadowaead stream.
type Cipher interface {
KeySize() int
SaltSize() int
Encrypter(salt []byte) (cipher.AEAD, error)
Decrypter(salt []byte) (cipher.AEAD, error)
}
const (
payloadSizeMask = 0x3FFF
bufSize = 17 * 1024
)
type aeadWriter struct {
io.Writer
cipher.AEAD
nonce [32]byte // should be sufficient for most nonce sizes
}
// newAEADWriter wraps an io.Writer with authenticated encryption.
func newAEADWriter(w io.Writer, aead cipher.AEAD) *aeadWriter {
return &aeadWriter{Writer: w, AEAD: aead}
}
// Write encrypts p and writes to the embedded io.Writer.
func (w *aeadWriter) Write(p []byte) (n int, err error) {
b := buf.Get(bufSize)
defer buf.Put(b)
nonce := w.nonce[:w.NonceSize()]
tag := w.Overhead()
off := 2 + tag
if len(p) == 0 {
b = b[:off]
b[0], b[1] = byte(0), byte(0)
w.Seal(b[:0], nonce, b[:2], nil)
increment(nonce)
_, err = w.Writer.Write(b)
return
}
for nr := 0; n < len(p) && err == nil; n += nr {
nr = payloadSizeMask
if n+nr > len(p) {
nr = len(p) - n
}
b = b[:off+nr+tag]
b[0], b[1] = byte(nr>>8), byte(nr)
w.Seal(b[:0], nonce, b[:2], nil)
increment(nonce)
w.Seal(b[:off], nonce, p[n:n+nr], nil)
increment(nonce)
_, err = w.Writer.Write(b)
}
return
}
type aeadReader struct {
io.Reader
cipher.AEAD
nonce [32]byte // should be sufficient for most nonce sizes
buf []byte // to be put back into bufPool
off int // offset to unconsumed part of buf
}
// newAEADReader wraps an io.Reader with authenticated decryption.
func newAEADReader(r io.Reader, aead cipher.AEAD) *aeadReader {
return &aeadReader{Reader: r, AEAD: aead}
}
// read and decrypt a record into p. len(p) >= max payload size + AEAD overhead.
func (r *aeadReader) read(p []byte) (int, error) {
nonce := r.nonce[:r.NonceSize()]
tag := r.Overhead()
p = p[:2+tag]
if _, err := io.ReadFull(r.Reader, p); err != nil {
return 0, err
}
_, err := r.Open(p[:0], nonce, p, nil)
increment(nonce)
if err != nil {
return 0, err
}
size := (int(p[0])<<8 + int(p[1])) & payloadSizeMask
if size == 0 {
return 0, ErrZeroChunk
}
p = p[:size+tag]
if _, err := io.ReadFull(r.Reader, p); err != nil {
return 0, err
}
_, err = r.Open(p[:0], nonce, p, nil)
increment(nonce)
if err != nil {
return 0, err
}
return size, nil
}
// Read reads from the embedded io.Reader, decrypts and writes to p.
func (r *aeadReader) Read(p []byte) (int, error) {
if r.buf == nil {
if len(p) >= payloadSizeMask+r.Overhead() {
return r.read(p)
}
b := buf.Get(bufSize)
n, err := r.read(b)
if err != nil {
buf.Put(b)
return 0, err
}
r.buf = b[:n]
r.off = 0
}
n := copy(p, r.buf[r.off:])
r.off += n
if r.off == len(r.buf) {
buf.Put(r.buf[:cap(r.buf)])
r.buf = nil
}
return n, nil
}
// increment little-endian encoded unsigned integer b. Wrap around on overflow.
func increment(b []byte) {
for i := range b {
b[i]++
if b[i] != 0 {
return
}
}
}
// aeadConn wraps a stream-oriented net.Conn with the shadowaead cipher.
type aeadConn struct {
net.Conn
Cipher
r *aeadReader
w *aeadWriter
}
// newAEADConn wraps a stream-oriented net.Conn with cipher.
func newAEADConn(c net.Conn, ciph Cipher) *aeadConn {
return &aeadConn{Conn: c, Cipher: ciph}
}
func (c *aeadConn) initReader() error {
salt := make([]byte, c.SaltSize())
if _, err := io.ReadFull(c.Conn, salt); err != nil {
return err
}
aead, err := c.Decrypter(salt)
if err != nil {
return err
}
c.r = newAEADReader(c.Conn, aead)
return nil
}
func (c *aeadConn) Read(b []byte) (int, error) {
if c.r == nil {
if err := c.initReader(); err != nil {
return 0, err
}
}
return c.r.Read(b)
}
func (c *aeadConn) initWriter() error {
salt := make([]byte, c.SaltSize())
if _, err := rand.Read(salt); err != nil {
return err
}
aead, err := c.Encrypter(salt)
if err != nil {
return err
}
_, err = c.Conn.Write(salt)
if err != nil {
return err
}
c.w = newAEADWriter(c.Conn, aead)
return nil
}
func (c *aeadConn) Write(b []byte) (int, error) {
if c.w == nil {
if err := c.initWriter(); err != nil {
return 0, err
}
}
return c.w.Write(b)
}
-408
View File
@@ -1,408 +0,0 @@
package snell
import (
"bytes"
"encoding/binary"
"errors"
"fmt"
"io"
"net"
"net/netip"
"sync"
"github.com/sagernet/sing/common/buf"
)
const (
Version1 = 1
Version2 = 2
Version3 = 3
Version4 = 4
Version5 = 5
DefaultSnellVersion = Version1
// max packet length
maxLength = 0x3FFF
)
const (
CommandPing byte = 0
CommandConnect byte = 1
CommandConnectV2 byte = 5
CommandUDP byte = 6
CommandUDPForward byte = 1
CommandTunnel byte = 0
CommandPong byte = 1
CommandError byte = 2
Version byte = 1
)
// Snell wraps an encrypted stream and handles the snell reply header.
type Snell struct {
net.Conn
buffer [1]byte
reply bool
}
func (s *Snell) Read(b []byte) (int, error) {
if err := s.ReadReply(); err != nil {
return 0, err
}
return s.Conn.Read(b)
}
func (s *Snell) ReadReply() error {
if s.reply {
return nil
}
if _, err := io.ReadFull(s.Conn, s.buffer[:]); err != nil {
return err
}
s.reply = true
if s.buffer[0] == CommandTunnel {
return nil
} else if s.buffer[0] != CommandError {
return errors.New("command not support")
}
if _, err := io.ReadFull(s.Conn, s.buffer[:]); err != nil {
return err
}
errcode := int(s.buffer[0])
if _, err := io.ReadFull(s.Conn, s.buffer[:]); err != nil {
return err
}
length := int(s.buffer[0])
msg := make([]byte, length)
if _, err := io.ReadFull(s.Conn, msg); err != nil {
return err
}
return fmt.Errorf("server reported code: %d, message: %s", errcode, string(msg))
}
func WriteHeader(conn net.Conn, host string, port uint, version int) error {
return WriteHeaderWithReuse(conn, host, port, version, false)
}
func WriteHeaderWithReuse(conn net.Conn, host string, port uint, version int, reuse bool) error {
buffer := &bytes.Buffer{}
buffer.WriteByte(Version)
if version == Version2 || reuse {
buffer.WriteByte(CommandConnectV2)
} else {
buffer.WriteByte(CommandConnect)
}
buffer.WriteByte(0)
buffer.WriteByte(uint8(len(host)))
buffer.WriteString(host)
binary.Write(buffer, binary.BigEndian, uint16(port))
if _, err := conn.Write(buffer.Bytes()); err != nil {
return err
}
return nil
}
func WriteUDPHeader(conn net.Conn, version int) error {
if version < Version3 {
return errors.New("unsupport UDP version")
}
_, err := conn.Write([]byte{Version, CommandUDP, 0x00})
return err
}
// HalfClose only works after the request negotiated the reuse command.
func HalfClose(conn net.Conn) error {
if err := writeZeroChunk(conn); err != nil {
return err
}
if s, ok := conn.(*Snell); ok {
s.reply = false
}
return nil
}
// StreamConn wraps a raw connection with the snell stream cipher for the given version.
func StreamConn(conn net.Conn, psk []byte, version int) *Snell {
if version >= Version4 {
return &Snell{Conn: newV4Conn(conn, psk)}
}
var cipher Cipher
if version != Version1 {
cipher = NewAES128GCM(psk)
} else {
cipher = NewChacha20Poly1305(psk)
}
return &Snell{Conn: newAEADConn(conn, cipher)}
}
// ServerStreamConn wraps a raw connection on the server side.
func ServerStreamConn(conn net.Conn, psk []byte, version int) *Snell {
stream := StreamConn(conn, psk, version)
stream.reply = true
return stream
}
func PacketConn(conn net.Conn) net.PacketConn {
return &packetConn{
Conn: conn,
}
}
func (s *Snell) WritePacketFrame(b []byte) (int, error) {
if fw, ok := s.Conn.(packetFrameWriter); ok {
return fw.WritePacketFrame(b)
}
return s.Conn.Write(b)
}
func WritePacket(w io.Writer, target, payload []byte) (int, error) {
maxPayloadLength := maxLength - udpRequestHeaderLength(target)
if maxPayloadLength <= 0 {
return 0, errors.New("snell UDP address too large")
}
if len(payload) <= maxPayloadLength {
return writePacket(w, target, payload)
}
return 0, errors.New("snell UDP payload too large")
}
func WritePacketResponse(w io.Writer, addr net.Addr, payload []byte) (int, error) {
buffer := &bytes.Buffer{}
target := parseAddrToSocksAddr(addr)
if len(target) == 0 {
return 0, errors.New("snell UDP response address invalid")
}
switch target[0] {
case atypIPv4:
if len(target) < 1+net.IPv4len+2 {
return 0, errors.New("snell UDP response address invalid")
}
buffer.WriteByte(0x04)
buffer.Write(target[1 : 1+net.IPv4len+2])
case atypIPv6:
if len(target) < 1+net.IPv6len+2 {
return 0, errors.New("snell UDP response address invalid")
}
buffer.WriteByte(0x06)
buffer.Write(target[1 : 1+net.IPv6len+2])
default:
return 0, errors.New("snell UDP response address invalid")
}
buffer.Write(payload)
var err error
if fw, ok := w.(packetFrameWriter); ok {
_, err = fw.WritePacketFrame(buffer.Bytes())
} else {
_, err = w.Write(buffer.Bytes())
}
if err != nil {
return 0, err
}
return len(payload), nil
}
// UDPRequest is a parsed snell UDP forward request.
type UDPRequest struct {
Host string
Ip netip.Addr
Port uint16
Payload []byte
}
func ParseUDPRequest(packet []byte) (UDPRequest, error) {
if len(packet) < 2 || packet[0] != CommandUDPForward {
return UDPRequest{}, errors.New("snell invalid UDP request")
}
if hostLen := int(packet[1]); hostLen != 0 {
if len(packet) <= 2+hostLen+2 {
return UDPRequest{}, errors.New("snell invalid UDP domain request")
}
offset := 2 + hostLen
return UDPRequest{
Host: string(packet[2:offset]),
Port: binary.BigEndian.Uint16(packet[offset : offset+2]),
Payload: packet[offset+2:],
}, nil
}
if len(packet) < 3 {
return UDPRequest{}, errors.New("snell invalid UDP IP request")
}
switch packet[2] {
case 0x04:
if len(packet) < 3+net.IPv4len+2 {
return UDPRequest{}, errors.New("snell invalid UDP IPv4 request")
}
offset := 3 + net.IPv4len
ip, _ := netip.AddrFromSlice(packet[3:offset])
return UDPRequest{
Ip: ip.Unmap(),
Port: binary.BigEndian.Uint16(packet[offset : offset+2]),
Payload: packet[offset+2:],
}, nil
case 0x06:
if len(packet) < 3+net.IPv6len+2 {
return UDPRequest{}, errors.New("snell invalid UDP IPv6 request")
}
offset := 3 + net.IPv6len
ip, _ := netip.AddrFromSlice(packet[3:offset])
return UDPRequest{
Ip: ip.Unmap(),
Port: binary.BigEndian.Uint16(packet[offset : offset+2]),
Payload: packet[offset+2:],
}, nil
default:
return UDPRequest{}, errors.New("snell invalid UDP address type")
}
}
func ReadPacket(r io.Reader, payload []byte) (net.Addr, int, error) {
b := buf.Get(buf.UDPBufferSize)
defer buf.Put(b)
n, err := r.Read(b)
headLen := 1
if err != nil {
return nil, 0, err
}
if n < headLen {
return nil, 0, errors.New("insufficient UDP length")
}
switch b[0] {
case 0x04:
headLen += net.IPv4len + 2
if n < headLen {
err = errors.New("insufficient UDP length")
break
}
b[0] = atypIPv4
case 0x06:
headLen += net.IPv6len + 2
if n < headLen {
err = errors.New("insufficient UDP length")
break
}
b[0] = atypIPv6
default:
err = errors.New("ip version invalid")
}
if err != nil {
return nil, 0, err
}
addr := splitSocksAddr(b[0:])
if addr == nil {
return nil, 0, errors.New("remote address invalid")
}
uAddr := addr.UDPAddr()
if uAddr == nil {
return nil, 0, errors.New("parse addr error")
}
length := len(payload)
if n-headLen < length {
length = n - headLen
}
copy(payload[:], b[headLen:headLen+length])
return uAddr, length, nil
}
var endSignal = []byte{}
type packetFrameWriter interface {
WritePacketFrame([]byte) (int, error)
}
func writeZeroChunk(conn net.Conn) error {
if _, err := conn.Write(endSignal); err != nil {
return err
}
return nil
}
func writePacket(w io.Writer, target, payload []byte) (int, error) {
buffer := &bytes.Buffer{}
buffer.WriteByte(CommandUDPForward)
switch target[0] {
case atypDomainName:
hostLen := target[1]
if len(target) < 1+1+int(hostLen)+2 {
return 0, errors.New("snell UDP address invalid")
}
buffer.Write(target[1 : 1+1+hostLen+2])
case atypIPv4:
if len(target) < 1+net.IPv4len+2 {
return 0, errors.New("snell UDP address invalid")
}
buffer.Write([]byte{0x00, 0x04})
buffer.Write(target[1 : 1+net.IPv4len+2])
case atypIPv6:
if len(target) < 1+net.IPv6len+2 {
return 0, errors.New("snell UDP address invalid")
}
buffer.Write([]byte{0x00, 0x06})
buffer.Write(target[1 : 1+net.IPv6len+2])
default:
return 0, errors.New("snell UDP address invalid")
}
buffer.Write(payload)
if fw, ok := w.(packetFrameWriter); ok {
_, err := fw.WritePacketFrame(buffer.Bytes())
if err != nil {
return 0, err
}
return len(payload), nil
}
_, err := w.Write(buffer.Bytes())
if err != nil {
return 0, err
}
return len(payload), nil
}
func udpRequestHeaderLength(target []byte) int {
if len(target) == 0 {
return maxLength + 1
}
switch target[0] {
case atypDomainName:
if len(target) < 2 {
return maxLength + 1
}
return 1 + 1 + int(target[1]) + 2
case atypIPv4:
return 1 + 2 + net.IPv4len + 2
case atypIPv6:
return 1 + 2 + net.IPv6len + 2
default:
return maxLength + 1
}
}
type packetConn struct {
net.Conn
rMux sync.Mutex
wMux sync.Mutex
}
func (pc *packetConn) WritePacketFrame(b []byte) (int, error) {
if s, ok := pc.Conn.(*Snell); ok {
if fw, ok := s.Conn.(packetFrameWriter); ok {
return fw.WritePacketFrame(b)
}
}
return pc.Conn.Write(b)
}
func (pc *packetConn) WriteTo(b []byte, addr net.Addr) (int, error) {
pc.wMux.Lock()
defer pc.wMux.Unlock()
return WritePacket(pc, parseAddrToSocksAddr(addr), b)
}
func (pc *packetConn) ReadFrom(b []byte) (int, net.Addr, error) {
pc.rMux.Lock()
defer pc.rMux.Unlock()
addr, n, err := ReadPacket(pc.Conn, b)
if err != nil {
return 0, nil, err
}
return n, addr, nil
}
-463
View File
@@ -1,463 +0,0 @@
package snell
import (
"crypto/cipher"
cryptorand "crypto/rand"
"encoding/binary"
"errors"
"io"
"math"
"math/big"
"math/bits"
"net"
"sync"
"time"
)
const (
v4SaltSize = 16
v4NonceSize = 12
v4HeaderPlainSize = 7
v4HeaderCipherSize = v4HeaderPlainSize + 16
v4FrameSize = 1460
v4InitialPaddingMin = 0x100
v4InitialPaddingSpan = 0x100
)
type v4Conn struct {
net.Conn
psk []byte
r *v4Reader
w *v4Writer
}
func newV4Conn(conn net.Conn, psk []byte) *v4Conn {
return &v4Conn{Conn: conn, psk: psk}
}
func (c *v4Conn) initReader() error {
salt := make([]byte, v4SaltSize)
if _, err := io.ReadFull(c.Conn, salt); err != nil {
return err
}
aead, err := v4AEAD(c.psk, salt)
if err != nil {
return err
}
c.r = &v4Reader{Reader: c.Conn, aead: aead}
return nil
}
func (c *v4Conn) initWriter() error {
w, err := newV4Writer(c.Conn, c.psk)
if err != nil {
return err
}
c.w = w
return nil
}
func (c *v4Conn) Read(b []byte) (int, error) {
if c.r == nil {
if err := c.initReader(); err != nil {
return 0, err
}
}
return c.r.Read(b)
}
func (c *v4Conn) Write(b []byte) (int, error) {
if c.w == nil {
if err := c.initWriter(); err != nil {
return 0, err
}
}
return c.w.Write(b)
}
func (c *v4Conn) WritePacketFrame(b []byte) (int, error) {
if len(b) > maxLength {
return 0, errors.New("snell v4 frame too large")
}
if c.w == nil {
if err := c.initWriter(); err != nil {
return 0, err
}
}
c.w.mux.Lock()
defer c.w.mux.Unlock()
if err := c.w.writeFrame(b, c.w.nextFramePaddingLength(len(b))); err != nil {
return 0, err
}
return len(b), nil
}
func (c *v4Conn) WriteTo(w io.Writer) (int64, error) {
if c.r == nil {
if err := c.initReader(); err != nil {
return 0, err
}
}
var written int64
buf := make([]byte, maxLength)
for {
n, err := c.r.Read(buf)
if n > 0 {
nw, ew := w.Write(buf[:n])
written += int64(nw)
if ew != nil {
return written, ew
}
if nw != n {
return written, io.ErrShortWrite
}
}
if err != nil {
if err == io.EOF {
err = nil
}
return written, err
}
}
}
func (c *v4Conn) ReadFrom(r io.Reader) (int64, error) {
if c.w == nil {
if err := c.initWriter(); err != nil {
return 0, err
}
}
var read int64
buf := make([]byte, maxLength)
for {
n, err := r.Read(buf)
if n > 0 {
read += int64(n)
if _, ew := c.w.Write(buf[:n]); ew != nil {
return read, ew
}
}
if err != nil {
if err == io.EOF {
err = nil
}
return read, err
}
}
}
func v4AEAD(psk, salt []byte) (cipher.AEAD, error) {
return aesGCM(snellKDF(psk, salt, 16))
}
type v4Reader struct {
io.Reader
aead cipher.AEAD
nonce [v4NonceSize]byte
buf []byte
mux sync.Mutex
}
func (r *v4Reader) Read(b []byte) (int, error) {
r.mux.Lock()
defer r.mux.Unlock()
if len(r.buf) == 0 {
payload, err := r.readFrame()
if err != nil {
return 0, err
}
r.buf = payload
}
n := copy(b, r.buf)
r.buf = r.buf[n:]
return n, nil
}
func (r *v4Reader) readFrame() ([]byte, error) {
headerCipher := make([]byte, v4HeaderCipherSize)
if _, err := io.ReadFull(r.Reader, headerCipher); err != nil {
return nil, err
}
header, err := r.aead.Open(headerCipher[:0], r.nonce[:], headerCipher, nil)
incrementV4Nonce(r.nonce[:])
if err != nil {
return nil, err
}
if len(header) != v4HeaderPlainSize || header[0] != 4 {
return nil, errors.New("snell v4 invalid frame header")
}
paddingLength := int(binary.BigEndian.Uint16(header[3:5]))
payloadLength := int(binary.BigEndian.Uint16(header[5:7]))
if payloadLength == 0 {
if paddingLength != 0 {
return nil, errors.New("snell v4 zero chunk with padding")
}
return nil, ErrZeroChunk
}
if payloadLength > maxLength || paddingLength > maxLength {
return nil, errors.New("snell v4 frame too large")
}
payloadCipherLength := payloadLength + r.aead.Overhead()
frame := make([]byte, paddingLength+payloadCipherLength)
if _, err := io.ReadFull(r.Reader, frame); err != nil {
return nil, err
}
if paddingLength > 0 {
swapPadding(frame[:paddingLength], frame[paddingLength:])
}
payloadCipher := frame[paddingLength:]
payload, err := r.aead.Open(payloadCipher[:0], r.nonce[:], payloadCipher, nil)
incrementV4Nonce(r.nonce[:])
if err != nil {
return nil, err
}
return payload, nil
}
type v4Writer struct {
io.Writer
aead cipher.AEAD
nonce [v4NonceSize]byte
salt [v4SaltSize]byte
saltSent bool
initialPaddingLength uint16
payloadLimit uint16
lastWrite time.Time
mux sync.Mutex
}
func newV4Writer(w io.Writer, psk []byte) (*v4Writer, error) {
var salt [v4SaltSize]byte
if _, err := io.ReadFull(cryptorand.Reader, salt[:]); err != nil {
return nil, err
}
aead, err := v4AEAD(psk, salt[:])
if err != nil {
return nil, err
}
paddingDelta, err := cryptoRandomInt(v4InitialPaddingSpan)
if err != nil {
return nil, err
}
return &v4Writer{
Writer: w,
aead: aead,
salt: salt,
initialPaddingLength: uint16(v4InitialPaddingMin + paddingDelta),
}, nil
}
func (w *v4Writer) Write(b []byte) (int, error) {
w.mux.Lock()
defer w.mux.Unlock()
if len(b) == 0 {
return 0, w.writeFrame(nil, 0)
}
written := 0
for written < len(b) {
payloadLimit := int(w.nextPayloadLimit())
if payloadLimit <= 0 || payloadLimit > maxLength {
payloadLimit = maxLength
}
end := written + payloadLimit
if end > len(b) {
end = len(b)
}
paddingLength := w.nextFramePaddingLength(end - written)
if err := w.writeFrame(b[written:end], paddingLength); err != nil {
return written, err
}
written = end
}
return written, nil
}
func (w *v4Writer) nextPayloadLimit() uint16 {
now := time.Now()
var payloadLimit uint16
switch {
case w.lastWrite.IsZero():
payloadLimit = v4FrameSize - 55 - w.initialPaddingLength
case now.Sub(w.lastWrite) > 30*time.Second:
payloadLimit = v4FrameSize - 39
default:
payloadLimit = w.payloadLimit
}
w.lastWrite = now
if payloadLimit <= maxLength-1 {
next := int(payloadLimit) + v4FrameSize - 39
if next > maxLength {
next = maxLength
}
w.payloadLimit = uint16(next)
} else {
w.payloadLimit = maxLength
}
return payloadLimit
}
func (w *v4Writer) nextFramePaddingLength(payloadLength int) int {
if w.saltSent || payloadLength == 0 {
return 0
}
return int(w.initialPaddingLength)
}
func (w *v4Writer) writeFrame(payload []byte, paddingLength int) error {
if len(payload) > maxLength || paddingLength > maxLength {
return errors.New("snell v4 frame too large")
}
if len(payload) == 0 && paddingLength != 0 {
return errors.New("snell v4 zero chunk with padding")
}
header := make([]byte, v4HeaderPlainSize)
header[0] = 4
binary.BigEndian.PutUint16(header[3:5], uint16(paddingLength))
binary.BigEndian.PutUint16(header[5:7], uint16(len(payload)))
headerCipher := w.aead.Seal(nil, w.nonce[:], header, nil)
incrementV4Nonce(w.nonce[:])
var payloadCipher []byte
if len(payload) > 0 {
payloadCipher = w.aead.Seal(nil, w.nonce[:], payload, nil)
incrementV4Nonce(w.nonce[:])
}
frameLength := len(headerCipher) + paddingLength + len(payloadCipher)
if !w.saltSent {
frameLength += v4SaltSize
}
frame := make([]byte, 0, frameLength)
if !w.saltSent {
frame = append(frame, w.salt[:]...)
w.saltSent = true
}
frame = append(frame, headerCipher...)
if paddingLength > 0 {
padding, err := makeV4Padding(payloadCipher, paddingLength)
if err != nil {
return err
}
swapPadding(padding, payloadCipher)
frame = append(frame, padding...)
}
frame = append(frame, payloadCipher...)
return writeFull(w.Writer, frame)
}
func swapPadding(padding, payloadCipher []byte) {
limit := len(padding)
if len(payloadCipher) < limit {
limit = len(payloadCipher)
}
for i := 0; i < limit; i += 2 {
padding[i], payloadCipher[i] = payloadCipher[i], padding[i]
}
}
func makeV4Padding(payloadCipher []byte, paddingLength int) ([]byte, error) {
if paddingLength <= 0 {
return nil, nil
}
payloadOnes := countV4PayloadOnes(payloadCipher)
payloadZeros := 8*len(payloadCipher) - payloadOnes
if payloadZeros <= 0 {
return makeV4RandomPadding(paddingLength)
}
ratio := float64(payloadOnes) / float64(payloadZeros)
if ratio <= 0.5 || ratio >= 1.6 {
return makeV4RandomPadding(paddingLength)
}
targetRatioBase := 1.6
if payloadZeros < payloadOnes {
targetRatioBase = 0.4
}
jitter, err := randomUnitFloat64()
if err != nil {
return nil, err
}
targetRatio := targetRatioBase + jitter/10
totalBits := 8 * (paddingLength + len(payloadCipher))
targetOnes := int(float64(totalBits)*(targetRatio/(targetRatio+1)) - float64(payloadOnes))
if targetOnes < 0 || targetOnes > 8*paddingLength {
return makeV4RandomPadding(paddingLength)
}
return makeV4BitCountPadding(paddingLength, targetOnes)
}
func countV4PayloadOnes(payloadCipher []byte) int {
limit := len(payloadCipher) &^ 3
ones := 0
for _, b := range payloadCipher[:limit] {
ones += bits.OnesCount8(b)
}
return ones
}
func makeV4RandomPadding(length int) ([]byte, error) {
padding := make([]byte, length)
_, err := io.ReadFull(cryptorand.Reader, padding)
return padding, err
}
func makeV4BitCountPadding(length, oneBits int) ([]byte, error) {
totalBits := 8 * length
if oneBits < 0 || oneBits > totalBits {
return nil, errors.New("snell v4 invalid padding bit count")
}
bitset := make([]byte, totalBits)
for i := 0; i < oneBits; i++ {
bitset[i] = 1
}
for i := totalBits - 1; i > 0; i-- {
j, err := cryptoRandomInt(i + 1)
if err != nil {
return nil, err
}
bitset[i], bitset[j] = bitset[j], bitset[i]
}
padding := make([]byte, length)
for i, bit := range bitset {
if bit == 1 {
padding[i/8] |= 1 << uint(i%8)
}
}
return padding, nil
}
func cryptoRandomInt(max int) (int, error) {
n, err := cryptorand.Int(cryptorand.Reader, big.NewInt(int64(max)))
if err != nil {
return 0, err
}
return int(n.Int64()), nil
}
func randomUnitFloat64() (float64, error) {
n, err := cryptorand.Int(cryptorand.Reader, big.NewInt(1<<53))
if err != nil {
return 0, err
}
return float64(n.Int64()) / math.Exp2(53), nil
}
func writeFull(w io.Writer, p []byte) error {
for len(p) > 0 {
n, err := w.Write(p)
if err != nil {
return err
}
if n == 0 {
return io.ErrShortWrite
}
p = p[n:]
}
return nil
}
func incrementV4Nonce(nonce []byte) {
for i := range nonce {
nonce[i]++
if nonce[i] != 0 {
return
}
}
}