Compare commits

..
2 Commits
Author SHA1 Message Date
Fangliding c84753dae6 fmt 2026-09-20 17:09:02 +08:00
Fangliding d17906c2f1 Optimize logger 2026-09-20 17:00:49 +08:00
135 changed files with 1431 additions and 6692 deletions
+1 -1
View File
@@ -23,7 +23,7 @@ func newFakeDNSSniffer(ctx context.Context) (protocolSnifferWithMetadata, error)
} }
if fakeDNSEngine == nil { if fakeDNSEngine == nil {
errNotInit := errors.New("FakeDNSEngine is not initialized, but such a sniffer is used").AtError() errNotInit := errors.New("FakeDNSEngine is not initialized, but such a sniffer is used")
return protocolSnifferWithMetadata{}, errNotInit return protocolSnifferWithMetadata{}, errNotInit
} }
return protocolSnifferWithMetadata{protocolSniffer: func(ctx context.Context, bytes []byte) (SniffResult, error) { return protocolSnifferWithMetadata{protocolSniffer: func(ctx context.Context, bytes []byte) (SniffResult, error) {
+1 -1
View File
@@ -28,7 +28,7 @@ func toNetIP(addrs []net.Address) ([]net.IP, error) {
if addr.Family().IsIP() { if addr.Family().IsIP() {
ips = append(ips, addr.IP()) ips = append(ips, addr.IP())
} else { } else {
return nil, errors.New("Failed to convert address", addr, "to Net IP.").AtWarning() return nil, errors.New("Failed to convert address", addr, "to Net IP.")
} }
} }
return ips, nil return ips, nil
+2 -2
View File
@@ -188,10 +188,10 @@ func parseResponse(payload []byte) (*IPRecord, error) {
var parser dnsmessage.Parser var parser dnsmessage.Parser
h, err := parser.Start(payload) h, err := parser.Start(payload)
if err != nil { if err != nil {
return nil, errors.New("failed to parse DNS response").Base(err).AtWarning() return nil, errors.New("failed to parse DNS response").Base(err)
} }
if err := parser.SkipAllQuestions(); err != nil { if err := parser.SkipAllQuestions(); err != nil {
return nil, errors.New("failed to skip questions in DNS response").Base(err).AtWarning() return nil, errors.New("failed to skip questions in DNS response").Base(err)
} }
now := time.Now() now := time.Now()
+3 -3
View File
@@ -58,7 +58,7 @@ func NewFakeDNSHolder() (*Holder, error) {
var err error var err error
if fkdns, err = NewFakeDNSHolderConfigOnly(nil); err != nil { if fkdns, err = NewFakeDNSHolderConfigOnly(nil); err != nil {
return nil, errors.New("Unable to create Fake Dns Engine").Base(err).AtError() return nil, errors.New("Unable to create Fake Dns Engine").Base(err)
} }
err = fkdns.initialize(dns.FakeIPv4Pool, 65535) err = fkdns.initialize(dns.FakeIPv4Pool, 65535)
if err != nil { if err != nil {
@@ -80,13 +80,13 @@ func (fkdns *Holder) initialize(ipPoolCidr string, lruSize int) error {
var err error var err error
if _, ipRange, err = net.ParseCIDR(ipPoolCidr); err != nil { if _, ipRange, err = net.ParseCIDR(ipPoolCidr); err != nil {
return errors.New("Unable to parse CIDR for Fake DNS IP assignment").Base(err).AtError() return errors.New("Unable to parse CIDR for Fake DNS IP assignment").Base(err)
} }
ones, bits := ipRange.Mask.Size() ones, bits := ipRange.Mask.Size()
rooms := bits - ones rooms := bits - ones
if math.Log2(float64(lruSize)) >= float64(rooms) { if math.Log2(float64(lruSize)) >= float64(rooms) {
return errors.New("LRU size is bigger than subnet size").AtError() return errors.New("LRU size is bigger than subnet size")
} }
fkdns.domainToIP = cache.NewLru(lruSize) fkdns.domainToIP = cache.NewLru(lruSize)
fkdns.ipRange = ipRange fkdns.ipRange = ipRange
+4 -4
View File
@@ -84,7 +84,7 @@ func NewServer(ctx context.Context, dest net.Destination, dispatcher routing.Dis
if dest.Network == net.Network_UDP { // UDP classic DNS mode if dest.Network == net.Network_UDP { // UDP classic DNS mode
return NewClassicNameServer(dest, dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP), nil return NewClassicNameServer(dest, dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP), nil
} }
return nil, errors.New("No available name server could be created from ", dest).AtWarning() return nil, errors.New("No available name server could be created from ", dest)
} }
// NewClient creates a DNS client managing a name server with client IP, domain rules and expected IPs. // NewClient creates a DNS client managing a name server with client IP, domain rules and expected IPs.
@@ -102,7 +102,7 @@ func NewClient(
// Create a new server for each client for now // Create a new server for each client for now
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP) server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
if err != nil { if err != nil {
return errors.New("failed to create nameserver").Base(err).AtWarning() return errors.New("failed to create nameserver").Base(err)
} }
_, isLocalDNS := server.(*LocalNameServer) _, isLocalDNS := server.(*LocalNameServer)
@@ -113,7 +113,7 @@ func NewClient(
if len(ns.ExpectedIp) > 0 { if len(ns.ExpectedIp) > 0 {
expectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.ExpectedIp) expectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.ExpectedIp)
if err != nil { if err != nil {
return errors.New("failed to create expected ip matcher").Base(err).AtWarning() return errors.New("failed to create expected ip matcher").Base(err)
} }
} }
@@ -122,7 +122,7 @@ func NewClient(
if len(ns.UnexpectedIp) > 0 { if len(ns.UnexpectedIp) > 0 {
unexpectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.UnexpectedIp) unexpectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.UnexpectedIp)
if err != nil { if err != nil {
return errors.New("failed to create unexpected ip matcher").Base(err).AtWarning() return errors.New("failed to create unexpected ip matcher").Base(err)
} }
} }
+2 -2
View File
@@ -27,7 +27,7 @@ func (s *FakeDNSServer) IsDisableCache() bool {
func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOption) ([]net.IP, uint32, error) { func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOption) ([]net.IP, uint32, error) {
if f.fakeDNSEngine == nil { if f.fakeDNSEngine == nil {
return nil, 0, errors.New("Unable to locate a fake DNS Engine").AtError() return nil, 0, errors.New("Unable to locate a fake DNS Engine")
} }
var ips []net.Address var ips []net.Address
@@ -39,7 +39,7 @@ func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOp
netIP, err := toNetIP(ips) netIP, err := toNetIP(ips)
if err != nil { if err != nil {
return nil, 0, errors.New("Unable to convert IP to net ip").Base(err).AtError() return nil, 0, errors.New("Unable to convert IP to net ip").Base(err)
} }
errors.LogInfo(ctx, f.Name(), " got answer: ", domain, " -> ", ips) errors.LogInfo(ctx, f.Name(), " got answer: ", domain, " -> ", ips)
+6 -2
View File
@@ -89,10 +89,10 @@ func (g *Instance) startInternal() error {
g.active = true g.active = true
if err := g.initAccessLogger(); err != nil { if err := g.initAccessLogger(); err != nil {
return errors.New("failed to initialize access logger").Base(err).AtWarning() return errors.New("failed to initialize access logger").Base(err)
} }
if err := g.initErrorLogger(); err != nil { if err := g.initErrorLogger(); err != nil {
return errors.New("failed to initialize error logger").Base(err).AtWarning() return errors.New("failed to initialize error logger").Base(err)
} }
return nil return nil
@@ -141,6 +141,10 @@ func (g *Instance) Handle(msg log.Message) {
} }
} }
func (g *Instance) Severity() log.Severity {
return g.config.ErrorLogLevel
}
// Close implements common.Closable.Close(). // Close implements common.Closable.Close().
func (g *Instance) Close() error { func (g *Instance) Close() error {
errors.LogDebug(context.Background(), "Logger closing") errors.LogDebug(context.Background(), "Logger closing")
+1 -1
View File
@@ -66,7 +66,7 @@ func NewAlwaysOnInboundHandler(ctx context.Context, tag string, receiverConfig *
} }
mss, err := internet.ToMemoryStreamConfig(receiverConfig.StreamSettings) mss, err := internet.ToMemoryStreamConfig(receiverConfig.StreamSettings)
if err != nil { if err != nil {
return nil, errors.New("failed to parse stream config").Base(err).AtWarning() return nil, errors.New("failed to parse stream config").Base(err)
} }
newCtx := session.ContextWithInbound(ctx, &session.Inbound{Tag: tag, Source: src}) newCtx := session.ContextWithInbound(ctx, &session.Inbound{Tag: tag, Source: src})
+1 -1
View File
@@ -165,7 +165,7 @@ func NewHandler(ctx context.Context, config *core.InboundHandlerConfig) (inbound
receiverSettings, ok := rawReceiverSettings.(*proxyman.ReceiverConfig) receiverSettings, ok := rawReceiverSettings.(*proxyman.ReceiverConfig)
if !ok { if !ok {
return nil, errors.New("not a ReceiverConfig").AtError() return nil, errors.New("not a ReceiverConfig")
} }
streamSettings := receiverSettings.StreamSettings streamSettings := receiverSettings.StreamSettings
+2 -2
View File
@@ -142,7 +142,7 @@ func (w *tcpWorker) Start() error {
go w.callback(conn) go w.callback(conn)
}) })
if err != nil { if err != nil {
return errors.New("failed to listen TCP on ", w.port).AtWarning().Base(err) return errors.New("failed to listen TCP on ", w.port).Base(err)
} }
w.hub = hub w.hub = hub
return nil return nil
@@ -528,7 +528,7 @@ func (w *dsWorker) Start() error {
go w.callback(conn) go w.callback(conn)
}) })
if err != nil { if err != nil {
return errors.New("failed to listen Unix Domain Socket on ", w.address).AtWarning().Base(err) return errors.New("failed to listen Unix Domain Socket on ", w.address).Base(err)
} }
w.hub = hub w.hub = hub
return nil return nil
+2 -2
View File
@@ -87,7 +87,7 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou
h.senderSettings = s h.senderSettings = s
mss, err := internet.ToMemoryStreamConfig(s.StreamSettings) mss, err := internet.ToMemoryStreamConfig(s.StreamSettings)
if err != nil { if err != nil {
return nil, errors.New("failed to parse stream settings").Base(err).AtWarning() return nil, errors.New("failed to parse stream settings").Base(err)
} }
h.streamSettings = mss h.streamSettings = mss
default: default:
@@ -217,7 +217,7 @@ func (h *Handler) Dispatch(ctx context.Context, link *transport.Link) {
if ob.Target.Network == net.Network_UDP && ob.Target.Port == 443 { if ob.Target.Network == net.Network_UDP && ob.Target.Port == 443 {
switch h.udp443 { switch h.udp443 {
case "reject": case "reject":
test(errors.New("XUDP rejected UDP/443 traffic").AtInfo()) test(errors.New("XUDP rejected UDP/443 traffic"))
return return
case "skip": case "skip":
goto out goto out
+2 -2
View File
@@ -68,13 +68,13 @@ func (p *Portal) HandleConnection(ctx context.Context, link *transport.Link) err
outbounds := session.OutboundsFromContext(ctx) outbounds := session.OutboundsFromContext(ctx)
ob := outbounds[len(outbounds)-1] ob := outbounds[len(outbounds)-1]
if ob == nil { if ob == nil {
return errors.New("outbound metadata not found").AtError() return errors.New("outbound metadata not found")
} }
if isDomain(ob.Target, p.domain) { if isDomain(ob.Target, p.domain) {
muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{}) muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{})
if err != nil { if err != nil {
return errors.New("failed to create mux client worker").Base(err).AtWarning() return errors.New("failed to create mux client worker").Base(err)
} }
worker, err := NewPortalWorker(muxClient) worker, err := NewPortalWorker(muxClient)
+2 -2
View File
@@ -115,7 +115,7 @@ func (rr *RoutingRule) BuildCondition() (Condition, error) {
} }
if conds.Len() == 0 { if conds.Len() == 0 {
return nil, errors.New("this rule has no effective fields").AtWarning() return nil, errors.New("this rule has no effective fields")
} }
return conds, nil return conds, nil
@@ -145,7 +145,7 @@ func (br *BalancingRule) Build(ohm outbound.Manager, dispatcher routing.Dispatch
} }
s, ok := i.(*StrategyLeastLoadConfig) s, ok := i.(*StrategyLeastLoadConfig)
if !ok { if !ok {
return nil, errors.New("not a StrategyLeastLoadConfig").AtError() return nil, errors.New("not a StrategyLeastLoadConfig")
} }
leastLoadStrategy := NewLeastLoadStrategy(s) leastLoadStrategy := NewLeastLoadStrategy(s)
return &Balancer{ return &Balancer{
+1 -11
View File
@@ -5,8 +5,7 @@ import (
) )
type windowsReader struct { type windowsReader struct {
bufs []syscall.WSABuf bufs []syscall.WSABuf
ready bool
} }
func (r *windowsReader) Init(bs []*Buffer) { func (r *windowsReader) Init(bs []*Buffer) {
@@ -16,7 +15,6 @@ func (r *windowsReader) Init(bs []*Buffer) {
for _, b := range bs { for _, b := range bs {
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]}) r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
} }
r.ready = false
} }
func (r *windowsReader) Clear() { func (r *windowsReader) Clear() {
@@ -27,14 +25,6 @@ func (r *windowsReader) Clear() {
} }
func (r *windowsReader) Read(fd uintptr) int32 { func (r *windowsReader) Read(fd uintptr) int32 {
// On the first invocation, we return -1 to indicate "not ready"
// to make rawConn.Read wait for readability using the runtime's own mechanism
// because syscall.WSARecv() is a blocking call when used with nil OVERLAPPED
if !r.ready {
r.ready = true
return -1
}
var nBytes uint32 var nBytes uint32
var flags uint32 var flags uint32
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil) err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
+13 -65
View File
@@ -18,17 +18,12 @@ type hasInnerError interface {
Unwrap() error Unwrap() error
} }
type hasSeverity interface {
Severity() log.Severity
}
// Error is an error object with underlying error. // Error is an error object with underlying error.
type Error struct { type Error struct {
prefix []interface{} prefix []interface{}
message []interface{} message []interface{}
caller string caller string
inner error inner error
severity log.Severity
} }
// Error implements error.Error(). // Error implements error.Error().
@@ -69,46 +64,6 @@ func (err *Error) Base(e error) *Error {
return err return err
} }
func (err *Error) atSeverity(s log.Severity) *Error {
err.severity = s
return err
}
func (err *Error) Severity() log.Severity {
if err.inner == nil {
return err.severity
}
if s, ok := err.inner.(hasSeverity); ok {
as := s.Severity()
if as < err.severity {
return as
}
}
return err.severity
}
// AtDebug sets the severity to debug.
func (err *Error) AtDebug() *Error {
return err.atSeverity(log.Severity_Debug)
}
// AtInfo sets the severity to info.
func (err *Error) AtInfo() *Error {
return err.atSeverity(log.Severity_Info)
}
// AtWarning sets the severity to warning.
func (err *Error) AtWarning() *Error {
return err.atSeverity(log.Severity_Warning)
}
// AtError sets the severity to error.
func (err *Error) AtError() *Error {
return err.atSeverity(log.Severity_Error)
}
// String returns the string representation of this error. // String returns the string representation of this error.
func (err *Error) String() string { func (err *Error) String() string {
return err.Error() return err.Error()
@@ -132,9 +87,8 @@ func New(msg ...interface{}) *Error {
details = details[:i] details = details[:i]
} }
return &Error{ return &Error{
message: msg, message: msg,
severity: log.Severity_Info, caller: details,
caller: details,
} }
} }
@@ -171,6 +125,9 @@ func LogErrorInner(ctx context.Context, inner error, msg ...interface{}) {
} }
func doLog(ctx context.Context, inner error, severity log.Severity, msg ...interface{}) { func doLog(ctx context.Context, inner error, severity log.Severity, msg ...interface{}) {
if log.GetSeverity() < severity {
return
}
pc, _, _, _ := runtime.Caller(2) pc, _, _, _ := runtime.Caller(2)
details := runtime.FuncForPC(pc).Name() details := runtime.FuncForPC(pc).Name()
if len(details) >= trim { if len(details) >= trim {
@@ -181,10 +138,9 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter
details = details[:i] details = details[:i]
} }
err := &Error{ err := &Error{
message: msg, message: msg,
severity: severity, caller: details,
caller: details, inner: inner,
inner: inner,
} }
if ctx != nil && ctx != context.Background() { if ctx != nil && ctx != context.Background() {
id := uint32(c.IDFromContext(ctx)) id := uint32(c.IDFromContext(ctx))
@@ -193,7 +149,7 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter
} }
} }
log.Record(&log.GeneralMessage{ log.Record(&log.GeneralMessage{
Severity: GetSeverity(err), Severity: severity,
Content: err, Content: err,
}) })
} }
@@ -217,11 +173,3 @@ L:
} }
return err return err
} }
// GetSeverity returns the actual severity of the error, including inner errors.
func GetSeverity(err error) log.Severity {
if s, ok := err.(hasSeverity); ok {
return s.Severity()
}
return log.Severity_Info
}
+6 -15
View File
@@ -7,30 +7,21 @@ import (
"github.com/google/go-cmp/cmp" "github.com/google/go-cmp/cmp"
. "github.com/xtls/xray-core/common/errors" . "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/log"
) )
func TestError(t *testing.T) { func TestError(t *testing.T) {
err := New("TestError") err := New("TestError")
if v := GetSeverity(err); v != log.Severity_Info { if v := err.Error(); !strings.Contains(v, "TestError") {
t.Error("severity: ", v) t.Error("error: ", v)
} }
err = New("TestError2").Base(io.EOF) err = New("TestError2").Base(io.EOF)
if v := GetSeverity(err); v != log.Severity_Info { if v := err.Error(); !strings.Contains(v, "EOF") {
t.Error("severity: ", v) t.Error("error: ", v)
} }
err = New("TestError3").Base(io.EOF).AtWarning() err = New("TestError3").Base(io.EOF)
if v := GetSeverity(err); v != log.Severity_Warning { err = New("TestError4").Base(err)
t.Error("severity: ", v)
}
err = New("TestError4").Base(io.EOF).AtWarning()
err = New("TestError5").Base(err)
if v := GetSeverity(err); v != log.Severity_Warning {
t.Error("severity: ", v)
}
if v := err.Error(); !strings.Contains(v, "EOF") { if v := err.Error(); !strings.Contains(v, "EOF") {
t.Error("error: ", v) t.Error("error: ", v)
} }
+21 -25
View File
@@ -1,7 +1,7 @@
package log // import "github.com/xtls/xray-core/common/log" package log // import "github.com/xtls/xray-core/common/log"
import ( import (
"sync" "sync/atomic"
"github.com/xtls/xray-core/common/serial" "github.com/xtls/xray-core/common/serial"
) )
@@ -29,36 +29,32 @@ func (m *GeneralMessage) String() string {
// Record writes a message into log stream. // Record writes a message into log stream.
func Record(msg Message) { func Record(msg Message) {
logHandler.Handle(msg) if h := logHandler.Load(); h != nil {
(*h).Handle(msg)
}
} }
var logHandler syncHandler type SeverityLogger interface {
Handler
Severity() Severity
}
func GetSeverity() Severity {
if h := logHandler.Load(); h != nil {
if sh, ok := (*h).(SeverityLogger); ok {
return sh.Severity()
}
}
// log everything by default
return Severity_Debug
}
var logHandler atomic.Pointer[Handler]
// RegisterHandler registers a new handler as current log handler. Previous registered handler will be discarded. // RegisterHandler registers a new handler as current log handler. Previous registered handler will be discarded.
func RegisterHandler(handler Handler) { func RegisterHandler(handler Handler) {
if handler == nil { if handler == nil {
panic("Log handler is nil") panic("Log handler is nil")
} }
logHandler.Set(handler) logHandler.Store(&handler)
}
type syncHandler struct {
sync.RWMutex
Handler
}
func (h *syncHandler) Handle(msg Message) {
h.RLock()
defer h.RUnlock()
if h.Handler != nil {
h.Handler.Handle(msg)
}
}
func (h *syncHandler) Set(handler Handler) {
h.Lock()
defer h.Unlock()
h.Handler = handler
} }
+4
View File
@@ -68,6 +68,10 @@ func (l *serverityLogger) Handle(msg Message) {
} }
} }
func (l *serverityLogger) Severity() Severity {
return l.logLevel
}
func (l *generalLogger) run() { func (l *generalLogger) run() {
defer l.access.Signal() defer l.access.Signal()
+1 -1
View File
@@ -38,7 +38,7 @@ func (m *ClientManager) Dispatch(ctx context.Context, link *transport.Link) erro
} }
} }
return errors.New("unable to find an available mux client").AtWarning() return errors.New("unable to find an available mux client")
} }
type WorkerPicker interface { type WorkerPicker interface {
+1 -1
View File
@@ -117,7 +117,7 @@ func (f *FrameMetadata) Unmarshal(reader io.Reader, readSourceAndLocal bool) err
return err return err
} }
if metaLen > 512 { if metaLen > 512 {
return errors.New("invalid metalen ", metaLen).AtError() return errors.New("invalid metalen ", metaLen)
} }
b := buf.New() b := buf.New()
+1 -1
View File
@@ -351,7 +351,7 @@ func (w *ServerWorker) handleFrame(ctx context.Context, reader *buf.BufferedRead
err = w.handleStatusKeep(&meta, reader) err = w.handleStatusKeep(&meta, reader)
default: default:
status := meta.SessionStatus status := meta.SessionStatus
return errors.New("unknown status: ", status).AtError() return errors.New("unknown status: ", status)
} }
if err != nil { if err != nil {
+1 -1
View File
@@ -7,7 +7,7 @@ import (
func (u *User) GetTypedAccount() (Account, error) { func (u *User) GetTypedAccount() (Account, error) {
if u.GetAccount() == nil { if u.GetAccount() == nil {
return nil, errors.New("Account is missing").AtWarning() return nil, errors.New("Account is missing")
} }
rawAccount, err := u.Account.GetInstance() rawAccount, err := u.Account.GetInstance()
+2 -2
View File
@@ -16,7 +16,7 @@ var typeCreatorRegistry = make(map[reflect.Type]ConfigCreator)
func RegisterConfig(config interface{}, configCreator ConfigCreator) error { func RegisterConfig(config interface{}, configCreator ConfigCreator) error {
configType := reflect.TypeOf(config) configType := reflect.TypeOf(config)
if _, found := typeCreatorRegistry[configType]; found { if _, found := typeCreatorRegistry[configType]; found {
return errors.New(configType.Name() + " is already registered").AtError() return errors.New(configType.Name() + " is already registered")
} }
typeCreatorRegistry[configType] = configCreator typeCreatorRegistry[configType] = configCreator
return nil return nil
@@ -27,7 +27,7 @@ func CreateObject(ctx context.Context, config interface{}) (interface{}, error)
configType := reflect.TypeOf(config) configType := reflect.TypeOf(config)
creator, found := typeCreatorRegistry[configType] creator, found := typeCreatorRegistry[configType]
if !found { if !found {
return nil, errors.New(configType.String() + " is not registered").AtError() return nil, errors.New(configType.String() + " is not registered")
} }
return creator(ctx, config) return creator(ctx, config)
} }
+4 -4
View File
@@ -125,7 +125,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
} }
if f == "" { if f == "" {
return nil, errors.New("Failed to get format of ", file).AtWarning() return nil, errors.New("Failed to get format of ", file)
} }
if f == "protobuf" { if f == "protobuf" {
@@ -142,7 +142,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
if len(v) == 1 { if len(v) == 1 {
return configLoaderByName["protobuf"].Loader(v) return configLoaderByName["protobuf"].Loader(v)
} else { } else {
return nil, errors.New("Only one protobuf config file is allowed").AtWarning() return nil, errors.New("Only one protobuf config file is allowed")
} }
} }
@@ -152,11 +152,11 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
if f, found := configLoaderByName[formatName]; found { if f, found := configLoaderByName[formatName]; found {
return f.Loader(v) return f.Loader(v)
} else { } else {
return nil, errors.New("Unable to load config in", formatName).AtWarning() return nil, errors.New("Unable to load config in", formatName)
} }
} }
return nil, errors.New("Unable to load config").AtWarning() return nil, errors.New("Unable to load config")
} }
func loadProtobufConfig(data []byte) (*Config, error) { func loadProtobufConfig(data []byte) (*Config, error) {
+5 -5
View File
@@ -24,11 +24,11 @@ require (
github.com/vishvananda/netlink v1.3.1 github.com/vishvananda/netlink v1.3.1
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
golang.org/x/crypto v0.57.0 golang.org/x/crypto v0.55.0
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
golang.org/x/net v0.59.0 golang.org/x/net v0.58.0
golang.org/x/sync v0.23.0 golang.org/x/sync v0.22.0
golang.org/x/sys v0.48.0 golang.org/x/sys v0.47.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
golang.zx2c4.com/wireguard/windows v1.0.1 golang.zx2c4.com/wireguard/windows v1.0.1
@@ -57,7 +57,7 @@ require (
github.com/vishvananda/netns v0.0.5 // indirect github.com/vishvananda/netns v0.0.5 // indirect
github.com/wlynxg/anet v0.0.5 // indirect github.com/wlynxg/anet v0.0.5 // indirect
go.yaml.in/yaml/v3 v3.0.5 // indirect go.yaml.in/yaml/v3 v3.0.5 // indirect
golang.org/x/text v0.42.0 // indirect golang.org/x/text v0.41.0 // indirect
golang.org/x/time v0.14.0 // indirect golang.org/x/time v0.14.0 // indirect
golang.org/x/tools v0.49.0 // indirect golang.org/x/tools v0.49.0 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
+10 -10
View File
@@ -111,8 +111,8 @@ go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBs
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y= go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M= golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA= golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM= golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc= golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY= golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
@@ -121,12 +121,12 @@ golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20211015210444-4f30a5c0130f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= golang.org/x/net v0.0.0-20211015210444-4f30a5c0130f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
golang.org/x/net v0.59.0 h1:5zfYln+w5XCxwrnMMJPufRgNoXEaGxl0wo5GqPXyues= golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
golang.org/x/net v0.59.0/go.mod h1:2DA/G1UfVbCpQPeWTmMPGY7Cs2PkBkwu743bVX5PIVg= golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -134,14 +134,14 @@ golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20211019181941-9d821ace8654/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20211019181941-9d821ace8654/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ= golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
golang.org/x/text v0.42.0 h1:JbOZXgfeCPU9gacVtYliJqOhD+zhrEqK4LfdpmlUZqI= golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
golang.org/x/text v0.42.0/go.mod h1:ojzP1Z+2QtioaF8DTtO8K5q7JWVVYwZKenzujK0Zd0E= golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI= golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4= golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
+2 -2
View File
@@ -97,7 +97,7 @@ func (v *HTTPClientConfig) Build() (proto.Message, error) {
user.Email = v.Email user.Email = v.Email
} else { } else {
if err := json.Unmarshal(rawUser, user); err != nil { if err := json.Unmarshal(rawUser, user); err != nil {
return nil, errors.New("failed to parse HTTP user").Base(err).AtError() return nil, errors.New("failed to parse HTTP user").Base(err)
} }
} }
account := new(HTTPAccount) account := new(HTTPAccount)
@@ -106,7 +106,7 @@ func (v *HTTPClientConfig) Build() (proto.Message, error) {
account.Password = v.Password account.Password = v.Password
} else { } else {
if err := json.Unmarshal(rawUser, account); err != nil { if err := json.Unmarshal(rawUser, account); err != nil {
return nil, errors.New("failed to parse HTTP account").Base(err).AtError() return nil, errors.New("failed to parse HTTP account").Base(err)
} }
} }
user.Account = serial.ToTypedMessage(account.Build()) user.Account = serial.ToTypedMessage(account.Build())
+1 -1
View File
@@ -18,7 +18,7 @@ func RegisterConfigureFilePostProcessingStage(name string, stage ConfigureFilePo
func PostProcessConfigureFile(conf *Config) error { func PostProcessConfigureFile(conf *Config) error {
for k, v := range configureFilePostProcessingStages { for k, v := range configureFilePostProcessingStages {
if err := v.Process(conf); err != nil { if err := v.Process(conf); err != nil {
return errors.New("Rejected by Postprocessing Stage ", k).AtError().Base(err) return errors.New("Rejected by Postprocessing Stage ", k).Base(err)
} }
} }
return nil return nil
+2 -2
View File
@@ -13,7 +13,7 @@ type ConfigCreatorCache map[string]ConfigCreator
func (v ConfigCreatorCache) RegisterCreator(id string, creator ConfigCreator) error { func (v ConfigCreatorCache) RegisterCreator(id string, creator ConfigCreator) error {
if _, found := v[id]; found { if _, found := v[id]; found {
return errors.New(id, " already registered.").AtError() return errors.New(id, " already registered.")
} }
v[id] = creator v[id] = creator
@@ -61,7 +61,7 @@ func (v *JSONConfigLoader) Load(raw []byte) (interface{}, string, error) {
} }
rawID, found := obj[v.idKey] rawID, found := obj[v.idKey]
if !found { if !found {
return nil, "", errors.New(v.idKey, " not found in JSON context").AtError() return nil, "", errors.New(v.idKey, " not found in JSON context")
} }
var id string var id string
if err := json.Unmarshal(rawID, &id); err != nil { if err := json.Unmarshal(rawID, &id); err != nil {
+1 -1
View File
@@ -30,7 +30,7 @@ func MergeConfigFromFiles(files []*core.ConfigSource) (string, error) {
if j, ok := creflect.MarshalToJson(c, true); ok { if j, ok := creflect.MarshalToJson(c, true); ok {
return j, nil return j, nil
} }
return "", errors.New("marshal to json failed.").AtError() return "", errors.New("marshal to json failed.")
} }
func mergeConfigs(files []*core.ConfigSource) (*conf.Config, error) { func mergeConfigs(files []*core.ConfigSource) (*conf.Config, error) {
+2 -3
View File
@@ -44,7 +44,6 @@ func (v *SocksServerConfig) Build() (proto.Message, error) {
case AuthMethodUserPass: case AuthMethodUserPass:
config.AuthType = socks.AuthType_PASSWORD config.AuthType = socks.AuthType_PASSWORD
default: default:
// errors.New("unknown socks auth method: ", v.AuthMethod, ". Default to noauth.").AtWarning().WriteToLog()
config.AuthType = socks.AuthType_NO_AUTH config.AuthType = socks.AuthType_NO_AUTH
} }
@@ -115,7 +114,7 @@ func (v *SocksClientConfig) Build() (proto.Message, error) {
user.Email = v.Email user.Email = v.Email
} else { } else {
if err := json.Unmarshal(rawUser, user); err != nil { if err := json.Unmarshal(rawUser, user); err != nil {
return nil, errors.New("failed to parse Socks user").Base(err).AtError() return nil, errors.New("failed to parse Socks user").Base(err)
} }
} }
account := new(SocksAccount) account := new(SocksAccount)
@@ -124,7 +123,7 @@ func (v *SocksClientConfig) Build() (proto.Message, error) {
account.Password = v.Password account.Password = v.Password
} else { } else {
if err := json.Unmarshal(rawUser, account); err != nil { if err := json.Unmarshal(rawUser, account); err != nil {
return nil, errors.New("failed to parse socks account").Base(err).AtError() return nil, errors.New("failed to parse socks account").Base(err)
} }
} }
user.Account = serial.ToTypedMessage(account.Build()) user.Account = serial.ToTypedMessage(account.Build())
+16 -5
View File
@@ -14,6 +14,7 @@ import (
googleuuid "github.com/google/uuid" googleuuid "github.com/google/uuid"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/finalmask/fragment" "github.com/xtls/xray-core/transport/internet/finalmask/fragment"
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom" "github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm" "github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
@@ -908,13 +909,22 @@ func (c *Realm) Build() (proto.Message, error) {
} }
type UDPHop struct { type UDPHop struct {
Mode string `json:"mode"` Sockopt *SocketConfig `json:"sockopt"`
Interval Int32Range `json:"interval"` Mode string `json:"mode"`
RemoteIPs []string `json:"remoteIPs"` Interval Int32Range `json:"interval"`
RemotePorts PortList `json:"remotePorts"` RemotePorts PortList `json:"remotePorts"`
RemoteIPs []string `json:"remoteIPs"`
} }
func (c *UDPHop) Build() (proto.Message, error) { func (c *UDPHop) Build() (proto.Message, error) {
var sockopt *internet.SocketConfig
if c.Sockopt != nil {
var err error
sockopt, err = c.Sockopt.Build()
if err != nil {
return nil, err
}
}
var local, remote, remoteOnce bool var local, remote, remoteOnce bool
for _, mode := range strings.Split(c.Mode, ",") { for _, mode := range strings.Split(c.Mode, ",") {
switch strings.ToLower(mode) { switch strings.ToLower(mode) {
@@ -943,13 +953,14 @@ func (c *UDPHop) Build() (proto.Message, error) {
return nil, errors.New("invalid ip ", ip) return nil, errors.New("invalid ip ", ip)
} }
return &udphop.Config{ return &udphop.Config{
Sockopt: sockopt,
Local: local, Local: local,
Remote: remote, Remote: remote,
RemoteOnce: remoteOnce, RemoteOnce: remoteOnce,
IntervalMin: int64(c.Interval.From), IntervalMin: int64(c.Interval.From),
IntervalMax: int64(c.Interval.To), IntervalMax: int64(c.Interval.To),
RemoteIPs: remoteIPs,
RemotePorts: c.RemotePorts.Build().Ports(), RemotePorts: c.RemotePorts.Build().Ports(),
RemoteIPs: remoteIPs,
}, nil }, nil
} }
-13
View File
@@ -36,8 +36,6 @@ func (p TransportProtocol) Build() (string, error) {
return "", errors.PrintRemovedFeatureError("QUIC transport (without web service, etc.)", "XHTTP stream-one H3") return "", errors.PrintRemovedFeatureError("QUIC transport (without web service, etc.)", "XHTTP stream-one H3")
case "hysteria": case "hysteria":
return "hysteria", nil return "hysteria", nil
case "xdrive":
return "xdrive", nil
default: default:
return "", errors.New("Config: unknown transport protocol: ", p) return "", errors.New("Config: unknown transport protocol: ", p)
} }
@@ -61,7 +59,6 @@ type StreamConfig struct {
WSSettings *WebSocketConfig `json:"wsSettings"` WSSettings *WebSocketConfig `json:"wsSettings"`
HTTPUPGRADESettings *HttpUpgradeConfig `json:"httpupgradeSettings"` HTTPUPGRADESettings *HttpUpgradeConfig `json:"httpupgradeSettings"`
HysteriaSettings *HysteriaConfig `json:"hysteriaSettings"` HysteriaSettings *HysteriaConfig `json:"hysteriaSettings"`
XDRIVESettings *XDriveConfig `json:"xdriveSettings"`
SocketSettings *SocketConfig `json:"sockopt"` SocketSettings *SocketConfig `json:"sockopt"`
} }
@@ -195,16 +192,6 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
Settings: serial.ToTypedMessage(hs), Settings: serial.ToTypedMessage(hs),
}) })
} }
if c.XDRIVESettings != nil {
xs, err := c.XDRIVESettings.Build()
if err != nil {
return nil, errors.New("Failed to build XDRIVE config.").Base(err)
}
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
ProtocolName: "xdrive",
Settings: serial.ToTypedMessage(xs),
})
}
if c.SocketSettings != nil { if c.SocketSettings != nil {
ss, err := c.SocketSettings.Build() ss, err := c.SocketSettings.Build()
if err != nil { if err != nil {
+4 -52
View File
@@ -23,7 +23,6 @@ import (
"github.com/xtls/xray-core/transport/internet/splithttp" "github.com/xtls/xray-core/transport/internet/splithttp"
"github.com/xtls/xray-core/transport/internet/tcp" "github.com/xtls/xray-core/transport/internet/tcp"
"github.com/xtls/xray-core/transport/internet/websocket" "github.com/xtls/xray-core/transport/internet/websocket"
"github.com/xtls/xray-core/transport/internet/xdrive"
"google.golang.org/protobuf/proto" "google.golang.org/protobuf/proto"
) )
@@ -122,7 +121,7 @@ func (v *AuthenticatorRequest) Build() (*http.RequestConfig, error) {
for _, key := range headerNames { for _, key := range headerNames {
value := v.Headers[key] value := v.Headers[key]
if value == nil { if value == nil {
return nil, errors.New("empty HTTP header value: " + key).AtError() return nil, errors.New("empty HTTP header value: " + key)
} }
config.Header = append(config.Header, &http.Header{ config.Header = append(config.Header, &http.Header{
Name: key, Name: key,
@@ -190,7 +189,7 @@ func (v *AuthenticatorResponse) Build() (*http.ResponseConfig, error) {
for _, key := range headerNames { for _, key := range headerNames {
value := v.Headers[key] value := v.Headers[key]
if value == nil { if value == nil {
return nil, errors.New("empty HTTP header value: " + key).AtError() return nil, errors.New("empty HTTP header value: " + key)
} }
config.Header = append(config.Header, &http.Header{ config.Header = append(config.Header, &http.Header{
Name: key, Name: key,
@@ -240,11 +239,11 @@ func (c *TCPConfig) Build() (proto.Message, error) {
if len(c.HeaderConfig) > 0 { if len(c.HeaderConfig) > 0 {
headerConfig, _, err := tcpHeaderLoader.Load(c.HeaderConfig) headerConfig, _, err := tcpHeaderLoader.Load(c.HeaderConfig)
if err != nil { if err != nil {
return nil, errors.New("invalid TCP header config").Base(err).AtError() return nil, errors.New("invalid TCP header config").Base(err)
} }
ts, err := headerConfig.(Buildable).Build() ts, err := headerConfig.(Buildable).Build()
if err != nil { if err != nil {
return nil, errors.New("invalid TCP header config").Base(err).AtError() return nil, errors.New("invalid TCP header config").Base(err)
} }
config.HeaderSettings = serial.ToTypedMessage(ts) config.HeaderSettings = serial.ToTypedMessage(ts)
} }
@@ -795,50 +794,3 @@ func readFileOrString(f string, s []string) ([]byte, error) {
} }
return nil, errors.New("both file and bytes are empty.") return nil, errors.New("both file and bytes are empty.")
} }
type XDriveConfig struct {
RemoteFolder string `json:"remoteFolder"`
Service string `json:"service"`
Secrets []string `json:"secrets"`
SegmentBytes uint32 `json:"segmentBytes"`
FlushIntervalMs uint32 `json:"flushIntervalMs"`
PollIntervalMs uint32 `json:"pollIntervalMs"`
MaxPollIntervalMs uint32 `json:"maxPollIntervalMs"`
SessionTTLSeconds uint32 `json:"sessionTtlSeconds"`
Concurrency uint32 `json:"concurrency"`
EagerWindowMs uint32 `json:"eagerWindowMs"`
HoleTimeoutMs uint32 `json:"holeTimeoutMs"`
Template json.RawMessage `json:"template"`
}
// Build implements Buildable.
func (c *XDriveConfig) Build() (proto.Message, error) {
switch c.Service {
case "local":
case "Google Drive":
if len(c.Secrets) != 3 {
return nil, errors.New("Google Drive needs 3 secrets in order of ClientID, ClientSecret, RefreshToken")
}
case "template":
if len(c.Template) == 0 {
return nil, errors.New(`service "template" needs a "template" object`)
}
default:
return nil, errors.New("unsupported service")
}
config := &xdrive.Config{
RemoteFolder: c.RemoteFolder,
Service: c.Service,
Secrets: c.Secrets,
SegmentBytes: c.SegmentBytes,
FlushIntervalMs: c.FlushIntervalMs,
PollIntervalMs: c.PollIntervalMs,
MaxPollIntervalMs: c.MaxPollIntervalMs,
SessionTtlSeconds: c.SessionTTLSeconds,
Concurrency: c.Concurrency,
EagerWindowMs: c.EagerWindowMs,
HoleTimeoutMs: c.HoleTimeoutMs,
Template: string(c.Template),
}
return config, nil
}
-73
View File
@@ -291,76 +291,3 @@ func TestHeaderCustomUDPBuildRejectsExprWithoutArgs(t *testing.T) {
t.Fatalf("expected transform arg rejection, got %v", err) t.Fatalf("expected transform arg rejection, got %v", err)
} }
} }
func TestXDriveStreamConfig(t *testing.T) {
config := new(StreamConfig)
if err := json.Unmarshal([]byte(`{
"method": "xdrive",
"xdriveSettings": {
"remoteFolder": "/tmp/xdrive",
"service": "local"
}
}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
built, err := config.Build()
if err != nil {
t.Fatalf("Build: %v", err)
}
if built.ProtocolName != "xdrive" {
t.Fatalf("ProtocolName is %q, want %q", built.ProtocolName, "xdrive")
}
if len(built.TransportSettings) != 1 || built.TransportSettings[0].ProtocolName != "xdrive" {
t.Fatalf("TransportSettings is %v, want a single xdrive entry", built.TransportSettings)
}
}
func TestXDriveRejectsUnknownService(t *testing.T) {
config := new(XDriveConfig)
if err := json.Unmarshal([]byte(`{"remoteFolder": "/tmp/xdrive", "service": "Dropbox"}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
if _, err := config.Build(); err == nil {
t.Fatal("Build accepted an unsupported service")
}
}
func TestXDriveTemplateStreamConfig(t *testing.T) {
config := new(StreamConfig)
if err := json.Unmarshal([]byte(`{
"method": "xdrive",
"xdriveSettings": {
"remoteFolder": "folder",
"service": "template",
"secrets": ["user", "pass"],
"template": {
"flatten": true,
"auth": {"type": "basic", "username": "{secret0}", "password": "{secret1}"},
"put": {"method": "PUT", "url": "https://dav.example/{folder}/{name}"},
"get": {"method": "GET", "url": "https://dav.example/{folder}/{name}"},
"delete": {"method": "DELETE", "url": "https://dav.example/{folder}/{name}"},
"list": {"method": "PROPFIND", "url": "https://dav.example/{folder}/", "namesRegex": "<d:href>/folder/([^<]+)</d:href>"}
}
}
}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
built, err := config.Build()
if err != nil {
t.Fatalf("Build: %v", err)
}
if built.ProtocolName != "xdrive" {
t.Fatalf("ProtocolName is %q, want xdrive", built.ProtocolName)
}
}
func TestXDriveTemplateNeedsTemplate(t *testing.T) {
config := new(XDriveConfig)
if err := json.Unmarshal([]byte(`{"remoteFolder": "f", "service": "template"}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
if _, err := config.Build(); err == nil {
t.Fatal("Build accepted a template service without a template")
}
}
+23 -7
View File
@@ -59,13 +59,14 @@ func (c *WireGuardPeerConfig) Build() (*wireguard.PeerConfig, error) {
type WireGuardConfig struct { type WireGuardConfig struct {
IsClient bool `json:""` IsClient bool `json:""`
NoKernelTun bool `json:"noKernelTun"` NoKernelTun bool `json:"noKernelTun"`
SecretKey string `json:"secretKey"` SecretKey string `json:"secretKey"`
Address []string `json:"address"` Address []string `json:"address"`
Peers []*WireGuardPeerConfig `json:"peers"` Peers []*WireGuardPeerConfig `json:"peers"`
MTU int32 `json:"mtu"` MTU int32 `json:"mtu"`
Reserved []byte `json:"reserved"` Reserved []byte `json:"reserved"`
DNS []string `json:"remoteDNS"` DomainStrategy string `json:"domainStrategy"`
DNS []string `json:"remoteDNS"`
} }
func (c *WireGuardConfig) Build() (proto.Message, error) { func (c *WireGuardConfig) Build() (proto.Message, error) {
@@ -124,6 +125,21 @@ func (c *WireGuardConfig) Build() (proto.Message, error) {
} }
config.Reserved = c.Reserved config.Reserved = c.Reserved
switch strings.ToLower(c.DomainStrategy) {
case "forceip", "":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP
case "forceipv4":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP4
case "forceipv6":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP6
case "forceipv4v6":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP46
case "forceipv6v4":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP64
default:
return nil, errors.New("unsupported domain strategy: ", c.DomainStrategy)
}
config.IsClient = c.IsClient config.IsClient = c.IsClient
config.NoKernelTun = c.NoKernelTun config.NoKernelTun = c.NoKernelTun
config.DNS = c.DNS config.DNS = c.DNS
-1
View File
@@ -60,7 +60,6 @@ import (
_ "github.com/xtls/xray-core/transport/internet/tls" _ "github.com/xtls/xray-core/transport/internet/tls"
_ "github.com/xtls/xray-core/transport/internet/udp" _ "github.com/xtls/xray-core/transport/internet/udp"
_ "github.com/xtls/xray-core/transport/internet/websocket" _ "github.com/xtls/xray-core/transport/internet/websocket"
_ "github.com/xtls/xray-core/transport/internet/xdrive"
// Transport headers // Transport headers
_ "github.com/xtls/xray-core/transport/internet/headers/http" _ "github.com/xtls/xray-core/transport/internet/headers/http"
+4 -8
View File
@@ -115,11 +115,7 @@ Start:
request, err := http.ReadRequest(reader) request, err := http.ReadRequest(reader)
if err != nil { if err != nil {
trace := errors.New("failed to read http request").Base(err) return errors.New("failed to read http request").Base(err)
if errors.Cause(err) != io.EOF && !isTimeout(errors.Cause(err)) {
trace.AtWarning()
}
return trace
} }
if len(s.config.Accounts) > 0 { if len(s.config.Accounts) > 0 {
@@ -147,7 +143,7 @@ Start:
} }
dest, err := http_proto.ParseHost(host, defaultPort) dest, err := http_proto.ParseHost(host, defaultPort)
if err != nil { if err != nil {
return errors.New("malformed proxy host: ", host).AtWarning().Base(err) return errors.New("malformed proxy host: ", host).Base(err)
} }
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{ ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(), From: conn.RemoteAddr(),
@@ -262,7 +258,7 @@ func (s *Server) handlePlainHTTP(ctx context.Context, request *http.Request, wri
requestWriter := buf.NewBufferedWriter(link.Writer) requestWriter := buf.NewBufferedWriter(link.Writer)
common.Must(requestWriter.SetBuffered(false)) common.Must(requestWriter.SetBuffered(false))
if err := request.Write(requestWriter); err != nil { if err := request.Write(requestWriter); err != nil {
return errors.New("failed to write whole request").Base(err).AtWarning() return errors.New("failed to write whole request").Base(err)
} }
return nil return nil
} }
@@ -299,7 +295,7 @@ func (s *Server) handlePlainHTTP(ctx context.Context, request *http.Request, wri
response.Header.Set("Proxy-Connection", "close") response.Header.Set("Proxy-Connection", "close")
} }
if err := response.Write(writer); err != nil { if err := response.Write(writer); err != nil {
return errors.New("failed to write response").Base(err).AtWarning() return errors.New("failed to write response").Base(err)
} }
return nil return nil
} }
+4 -4
View File
@@ -62,7 +62,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
conn, err := dialer.Dial(hysteria.ContextWithDatagram(ctx, target.Network == net.Network_UDP), c.server.Destination) conn, err := dialer.Dial(hysteria.ContextWithDatagram(ctx, target.Network == net.Network_UDP), c.server.Destination)
if err != nil { if err != nil {
return errors.New("failed to find an available destination").AtWarning().Base(err) return errors.New("failed to find an available destination").Base(err)
} }
defer conn.Close() defer conn.Close()
errors.LogInfo(ctx, "tunneling request to ", target, " via ", target.Network, ":", c.server.Destination.NetAddr()) errors.LogInfo(ctx, "tunneling request to ", target, " via ", target.Network, ":", c.server.Destination.NetAddr())
@@ -236,14 +236,14 @@ type UDPReader struct {
func (r *UDPReader) ReadFrom(p []byte) (n int, addr *net.Destination, err error) { func (r *UDPReader) ReadFrom(p []byte) (n int, addr *net.Destination, err error) {
for { for {
var packet [1500]byte var buf [hysteria.MaxDatagramFrameSize]byte
n, err := r.reader.Read(packet[:]) n, err := r.reader.Read(buf[:])
if err != nil { if err != nil {
return 0, nil, err return 0, nil, err
} }
msg, err := ParseUDPMessage(packet[:n]) msg, err := ParseUDPMessage(buf[:n])
if err != nil { if err != nil {
continue continue
} }
+2 -2
View File
@@ -40,11 +40,11 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) {
for _, user := range config.Users { for _, user := range config.Users {
u, err := user.ToMemoryUser() u, err := user.ToMemoryUser()
if err != nil { if err != nil {
return nil, errors.New("failed to get hysteria user").Base(err).AtError() return nil, errors.New("failed to get hysteria user").Base(err)
} }
if err := validator.Add(u); err != nil { if err := validator.Add(u); err != nil {
return nil, errors.New("failed to add user").Base(err).AtError() return nil, errors.New("failed to add user").Base(err)
} }
} }
+1 -1
View File
@@ -56,7 +56,7 @@ func (l *Loopback) init(config *Config, dispatcherInstance routing.Dispatcher) e
if config.Sniffing.GetEnabled() { if config.Sniffing.GetEnabled() {
request, err := proxyman.BuildSniffingRequest(config.Sniffing) request, err := proxyman.BuildSniffingRequest(config.Sniffing)
if err != nil { if err != nil {
return errors.New("failed to build loopback sniffing request").Base(err).AtError() return errors.New("failed to build loopback sniffing request").Base(err)
} }
l.sniffingRequest = request l.sniffingRequest = request
} }
-15
View File
@@ -277,7 +277,6 @@ func (w *VisionReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
w.ob.CanSpliceCopy = 1 w.ob.CanSpliceCopy = 1
} }
} }
SuppressOuterCloseNotify(w.conn)
readerConn, readCounter, _ := UnwrapRawConn(w.conn) readerConn, readCounter, _ := UnwrapRawConn(w.conn)
w.directReadCounter = readCounter w.directReadCounter = readCounter
w.Reader = buf.NewReader(readerConn) w.Reader = buf.NewReader(readerConn)
@@ -341,7 +340,6 @@ func (w *VisionWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
// w.ob.CanSpliceCopy = 1 // w.ob.CanSpliceCopy = 1
// } // }
} }
SuppressOuterCloseNotify(w.conn)
rawConn, _, writerCounter := UnwrapRawConn(w.conn) rawConn, _, writerCounter := UnwrapRawConn(w.conn)
w.Writer = buf.NewWriter(rawConn) w.Writer = buf.NewWriter(rawConn)
w.directWriteCounter = writerCounter w.directWriteCounter = writerCounter
@@ -671,19 +669,6 @@ func XtlsFilterTls(buffer buf.MultiBuffer, trafficState *TrafficState, ctx conte
} }
} }
type CloseNotifySuppressor interface {
SuppressCloseNotify()
}
// Close our local TLS conn instance might send a incorrect close_notify alert
// if the XTLS direct copy mode is enabled and cause TLS BAD_RECORD_MAC on users' browser
// Close the underlying connection directly to avoid this issue.
func SuppressOuterCloseNotify(conn net.Conn) {
if suppressor, ok := stat.TryUnwrapStatsConn(conn).(CloseNotifySuppressor); ok {
suppressor.SuppressCloseNotify()
}
}
// UnwrapRawConn support unwrap encryption, stats, mask wrappers, tls, utls, reality, proxyproto, uds-wrapper conn and get raw tcp/uds conn from it // UnwrapRawConn support unwrap encryption, stats, mask wrappers, tls, utls, reality, proxyproto, uds-wrapper conn and get raw tcp/uds conn from it
func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) { func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) {
var readCounter, writerCounter stats.Counter var readCounter, writerCounter stats.Counter
+2 -2
View File
@@ -71,7 +71,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
return nil return nil
}) })
if err != nil { if err != nil {
return errors.New("failed to find an available destination").AtWarning().Base(err) return errors.New("failed to find an available destination").Base(err)
} }
errors.LogInfo(ctx, "tunneling request to ", destination, " via ", network, ":", server.Destination.NetAddr()) errors.LogInfo(ctx, "tunneling request to ", destination, " via ", network, ":", server.Destination.NetAddr())
@@ -124,7 +124,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
} }
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout { if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
return errors.New("failed to write A request payload").Base(err).AtWarning() return errors.New("failed to write A request payload").Base(err)
} }
if err := bufferedWriter.SetBuffered(false); err != nil { if err := bufferedWriter.SetBuffered(false); err != nil {
+2 -2
View File
@@ -98,7 +98,7 @@ func ReadTCPSession(validator *Validator, reader io.Reader) (*protocol.RequestHe
iv := append([]byte(nil), buffer.BytesTo(ivLen)...) iv := append([]byte(nil), buffer.BytesTo(ivLen)...)
r, err = account.Cipher.NewDecryptionReader(account.Key, iv, reader) r, err = account.Cipher.NewDecryptionReader(account.Key, iv, reader)
if err != nil { if err != nil {
return nil, nil, drain.WithError(drainer, reader, errors.New("failed to initialize decoding stream").Base(err).AtError()) return nil, nil, drain.WithError(drainer, reader, errors.New("failed to initialize decoding stream").Base(err))
} }
} }
} }
@@ -146,7 +146,7 @@ func WriteTCPRequest(request *protocol.RequestHeader, writer io.Writer) (buf.Wri
w, err := account.Cipher.NewEncryptionWriter(account.Key, iv, writer) w, err := account.Cipher.NewEncryptionWriter(account.Key, iv, writer)
if err != nil { if err != nil {
return nil, errors.New("failed to create encoding stream").Base(err).AtError() return nil, errors.New("failed to create encoding stream").Base(err)
} }
header := buf.New() header := buf.New()
+3 -3
View File
@@ -34,11 +34,11 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) {
for _, user := range config.Users { for _, user := range config.Users {
u, err := user.ToMemoryUser() u, err := user.ToMemoryUser()
if err != nil { if err != nil {
return nil, errors.New("failed to get shadowsocks user").Base(err).AtError() return nil, errors.New("failed to get shadowsocks user").Base(err)
} }
if err := validator.Add(u); err != nil { if err := validator.Add(u); err != nil {
return nil, errors.New("failed to add user").Base(err).AtError() return nil, errors.New("failed to add user").Base(err)
} }
} }
@@ -200,7 +200,7 @@ func (s *Server) handleUDPPayload(ctx context.Context, conn stat.Connection, dis
func (s *Server) handleConnection(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { func (s *Server) handleConnection(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
sessionPolicy := s.policyManager.ForLevel(0) sessionPolicy := s.policyManager.ForLevel(0)
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil { if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
return errors.New("unable to set read deadline").Base(err).AtWarning() return errors.New("unable to set read deadline").Base(err)
} }
bufferedReader := buf.BufferedReader{Reader: buf.NewReader(conn)} bufferedReader := buf.BufferedReader{Reader: buf.NewReader(conn)}
+1 -1
View File
@@ -59,7 +59,7 @@ func NewMultiServer(ctx context.Context, config *MultiUserServerConfig) (*MultiU
} }
u, err := user.ToMemoryUser() u, err := user.ToMemoryUser()
if err != nil { if err != nil {
return nil, errors.New("failed to get shadowsocks user").Base(err).AtError() return nil, errors.New("failed to get shadowsocks user").Base(err)
} }
memUsers = append(memUsers, u) memUsers = append(memUsers, u)
} }
+1 -1
View File
@@ -105,7 +105,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
} }
udpRequest, err := ClientHandshake(request, conn, conn) udpRequest, err := ClientHandshake(request, conn, conn)
if err != nil { if err != nil {
return errors.New("failed to establish connection to server").AtWarning().Base(err) return errors.New("failed to establish connection to server").Base(err)
} }
if udpRequest != nil { if udpRequest != nil {
if udpRequest.Address == net.AnyIP || udpRequest.Address == net.AnyIPv6 { if udpRequest.Address == net.AnyIP || udpRequest.Address == net.AnyIPv6 {
+2 -2
View File
@@ -458,10 +458,10 @@ func ClientHandshake(request *protocol.RequestHeader, reader io.Reader, writer i
} }
if b.Byte(0) != socks5Version { if b.Byte(0) != socks5Version {
return nil, errors.New("unexpected server version: ", b.Byte(0)).AtWarning() return nil, errors.New("unexpected server version: ", b.Byte(0))
} }
if b.Byte(1) != authByte { if b.Byte(1) != authByte {
return nil, errors.New("auth method not supported.").AtWarning() return nil, errors.New("auth method not supported.")
} }
if authByte == authPassword { if authByte == authPassword {
+5 -5
View File
@@ -69,7 +69,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
return nil return nil
}) })
if err != nil { if err != nil {
return errors.New("failed to find an available destination").AtWarning().Base(err) return errors.New("failed to find an available destination").Base(err)
} }
errors.LogInfo(ctx, "tunneling request to ", destination, " via ", server.Destination.NetAddr()) errors.LogInfo(ctx, "tunneling request to ", destination, " via ", server.Destination.NetAddr())
@@ -116,21 +116,21 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
// write some request payload to buffer // write some request payload to buffer
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout { if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
return errors.New("failed to write A request payload").Base(err).AtWarning() return errors.New("failed to write A request payload").Base(err)
} }
// Flush; bufferWriter.WriteMultiBuffer now is bufferWriter.writer.WriteMultiBuffer // Flush; bufferWriter.WriteMultiBuffer now is bufferWriter.writer.WriteMultiBuffer
if err = bufferWriter.SetBuffered(false); err != nil { if err = bufferWriter.SetBuffered(false); err != nil {
return errors.New("failed to flush payload").Base(err).AtWarning() return errors.New("failed to flush payload").Base(err)
} }
// Send header if not sent yet // Send header if not sent yet
if _, err = connWriter.Write([]byte{}); err != nil { if _, err = connWriter.Write([]byte{}); err != nil {
return err.(*errors.Error).AtWarning() return err
} }
if err = buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer)); err != nil { if err = buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer)); err != nil {
return errors.New("failed to transfer request payload").Base(err).AtInfo() return errors.New("failed to transfer request payload").Base(err)
} }
return nil return nil
+12 -12
View File
@@ -47,11 +47,11 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) {
for _, user := range config.Users { for _, user := range config.Users {
u, err := user.ToMemoryUser() u, err := user.ToMemoryUser()
if err != nil { if err != nil {
return nil, errors.New("failed to get trojan user").Base(err).AtError() return nil, errors.New("failed to get trojan user").Base(err)
} }
if err := validator.Add(u); err != nil { if err := validator.Add(u); err != nil {
return nil, errors.New("failed to add user").Base(err).AtError() return nil, errors.New("failed to add user").Base(err)
} }
} }
@@ -151,7 +151,7 @@ func (s *Server) Process(ctx context.Context, network net.Network, conn stat.Con
sessionPolicy := s.policyManager.ForLevel(0) sessionPolicy := s.policyManager.ForLevel(0)
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil { if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
return errors.New("unable to set read deadline").Base(err).AtWarning() return errors.New("unable to set read deadline").Base(err)
} }
first := buf.FromBytes(make([]byte, buf.Size)) first := buf.FromBytes(make([]byte, buf.Size))
@@ -219,7 +219,7 @@ func (s *Server) Process(ctx context.Context, network net.Network, conn stat.Con
destination := clientReader.Target destination := clientReader.Target
if err := conn.SetReadDeadline(time.Time{}); err != nil { if err := conn.SetReadDeadline(time.Time{}); err != nil {
return errors.New("unable to set read deadline").Base(err).AtWarning() return errors.New("unable to set read deadline").Base(err)
} }
inbound := session.InboundFromContext(ctx) inbound := session.InboundFromContext(ctx)
@@ -402,7 +402,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
} }
apfb := napfb[name] apfb := napfb[name]
if apfb == nil { if apfb == nil {
return errors.New(`failed to find the default "name" config`).AtWarning() return errors.New(`failed to find the default "name" config`)
} }
if apfb[alpn] == nil { if apfb[alpn] == nil {
@@ -410,7 +410,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
} }
pfb := apfb[alpn] pfb := apfb[alpn]
if pfb == nil { if pfb == nil {
return errors.New(`failed to find the default "alpn" config`).AtWarning() return errors.New(`failed to find the default "alpn" config`)
} }
path := "" path := ""
@@ -444,7 +444,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
} }
fb := pfb[path] fb := pfb[path]
if fb == nil { if fb == nil {
return errors.New(`failed to find the default "path" config`).AtWarning() return errors.New(`failed to find the default "path" config`)
} }
ctx, cancel := context.WithCancel(ctx) ctx, cancel := context.WithCancel(ctx)
@@ -460,7 +460,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
} }
return nil return nil
}); err != nil { }); err != nil {
return errors.New("failed to dial to " + fb.Dest).Base(err).AtWarning() return errors.New("failed to dial to " + fb.Dest).Base(err)
} }
defer conn.Close() defer conn.Close()
@@ -520,11 +520,11 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
common.Must2(pro.Write([]byte{byte(p1 >> 8), byte(p1), byte(p2 >> 8), byte(p2)})) common.Must2(pro.Write([]byte{byte(p1 >> 8), byte(p1), byte(p2 >> 8), byte(p2)}))
} }
if err := serverWriter.WriteMultiBuffer(buf.MultiBuffer{pro}); err != nil { if err := serverWriter.WriteMultiBuffer(buf.MultiBuffer{pro}); err != nil {
return errors.New("failed to set PROXY protocol v", fb.Xver).Base(err).AtWarning() return errors.New("failed to set PROXY protocol v", fb.Xver).Base(err)
} }
} }
if err := buf.Copy(reader, serverWriter, buf.UpdateActivity(timer)); err != nil { if err := buf.Copy(reader, serverWriter, buf.UpdateActivity(timer)); err != nil {
return errors.New("failed to fallback request payload").Base(err).AtInfo() return errors.New("failed to fallback request payload").Base(err)
} }
return nil return nil
} }
@@ -534,7 +534,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
getResponse := func() error { getResponse := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
if err := buf.Copy(serverReader, writer, buf.UpdateActivity(timer)); err != nil { if err := buf.Copy(serverReader, writer, buf.UpdateActivity(timer)); err != nil {
return errors.New("failed to deliver response payload").Base(err).AtInfo() return errors.New("failed to deliver response payload").Base(err)
} }
return nil return nil
} }
@@ -542,7 +542,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
if err := task.Run(ctx, task.OnSuccess(postRequest, task.Close(serverWriter)), task.OnSuccess(getResponse, task.Close(writer))); err != nil { if err := task.Run(ctx, task.OnSuccess(postRequest, task.Close(serverWriter)), task.OnSuccess(getResponse, task.Close(writer))); err != nil {
common.Must(common.Interrupt(serverReader)) common.Must(common.Interrupt(serverReader))
common.Must(common.Interrupt(serverWriter)) common.Must(common.Interrupt(serverWriter))
return errors.New("fallback ends").Base(err).AtInfo() return errors.New("fallback ends").Base(err)
} }
return nil return nil
+4 -33
View File
@@ -37,25 +37,6 @@ type Handler struct {
downlinkCounter stats.Counter downlinkCounter stats.Counter
} }
type tunUDPStatsWriter struct {
writer buf.Writer
counter stats.Counter
}
func (w *tunUDPStatsWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
for len(mb) > 0 {
remaining, packet := buf.SplitFirst(mb)
packetSize := packet.Len()
if err := w.writer.WriteMultiBuffer(buf.MultiBuffer{packet}); err != nil {
buf.ReleaseMulti(remaining)
return err
}
w.counter.Add(int64(packetSize))
mb = remaining
}
return nil
}
// ConnectionHandler interface with the only method that stack is going to push new connections to // ConnectionHandler interface with the only method that stack is going to push new connections to
type ConnectionHandler interface { type ConnectionHandler interface {
HandleConnection(conn net.Conn, destination net.Destination) HandleConnection(conn net.Conn, destination net.Destination)
@@ -123,7 +104,7 @@ func (t *Handler) Start() error {
iface := updater.Get() iface := updater.Get()
if iface == nil { if iface == nil {
errors.LogInfo(context.Background(), "[tun] falied to set interface > iface == nil") errors.LogInfo(context.Background(), "[tun] falied to set interface > iface == nil")
return errors.New("iface not found") return nil
} }
return c.Control(func(fd uintptr) { return c.Control(func(fd uintptr) {
addrPort, _ := netip.ParseAddrPort(address) addrPort, _ := netip.ParseAddrPort(address)
@@ -190,8 +171,7 @@ func (t *Handler) HandleConnection(conn net.Conn, destination net.Destination) {
return return
} }
source := net.DestinationFromAddr(remote) source := net.DestinationFromAddr(remote)
isUDP := destination.Network == net.Network_UDP if t.uplinkCounter != nil || t.downlinkCounter != nil {
if !isUDP && (t.uplinkCounter != nil || t.downlinkCounter != nil) {
conn = &stat.CounterConnection{ conn = &stat.CounterConnection{
Connection: conn, Connection: conn,
ReadCounter: t.uplinkCounter, ReadCounter: t.uplinkCounter,
@@ -223,18 +203,9 @@ func (t *Handler) HandleConnection(conn net.Conn, destination net.Destination) {
}) })
errors.LogInfo(ctx, "processing from ", source, " to ", destination) errors.LogInfo(ctx, "processing from ", source, " to ", destination)
reader := &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)}
writer := buf.NewWriter(conn)
if isUDP {
reader.Counter = t.uplinkCounter
if t.downlinkCounter != nil {
writer = &tunUDPStatsWriter{writer: writer, counter: t.downlinkCounter}
}
}
link := &transport.Link{ link := &transport.Link{
Reader: reader, Reader: &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)},
Writer: writer, Writer: buf.NewWriter(conn),
} }
if err := t.dispatcher.DispatchLink(ctx, destination, link); err != nil { if err := t.dispatcher.DispatchLink(ctx, destination, link); err != nil {
errors.LogError(ctx, errors.New("connection closed").Base(err)) errors.LogError(ctx, errors.New("connection closed").Base(err))
+1 -1
View File
@@ -12,7 +12,7 @@ import (
func (a *Account) AsAccount() (protocol.Account, error) { func (a *Account) AsAccount() (protocol.Account, error) {
id, err := uuid.ParseString(a.Id) id, err := uuid.ParseString(a.Id)
if err != nil { if err != nil {
return nil, errors.New("failed to parse ID").Base(err).AtError() return nil, errors.New("failed to parse ID").Base(err)
} }
return &MemoryAccount{ return &MemoryAccount{
ID: protocol.NewID(id), ID: protocol.NewID(id),
+25 -25
View File
@@ -61,10 +61,10 @@ func init() {
for _, user := range c.Users { for _, user := range c.Users {
u, err := user.ToMemoryUser() u, err := user.ToMemoryUser()
if err != nil { if err != nil {
return nil, errors.New("failed to get VLESS user").Base(err).AtError() return nil, errors.New("failed to get VLESS user").Base(err)
} }
if err := validator.Add(u); err != nil { if err := validator.Add(u); err != nil {
return nil, errors.New("failed to initiate user").Base(err).AtError() return nil, errors.New("failed to initiate user").Base(err)
} }
} }
@@ -110,7 +110,7 @@ func New(ctx context.Context, config *Config, dc dns.Client, validator vless.Val
} }
handler.decryption = &encryption.ServerInstance{} handler.decryption = &encryption.ServerInstance{}
if err := handler.decryption.Init(nfsSKeysBytes, config.XorMode, config.SecondsFrom, config.SecondsTo, config.Padding); err != nil { if err := handler.decryption.Init(nfsSKeysBytes, config.XorMode, config.SecondsFrom, config.SecondsTo, config.Padding); err != nil {
return nil, errors.New("failed to use decryption").Base(err).AtError() return nil, errors.New("failed to use decryption").Base(err)
} }
} }
@@ -128,7 +128,7 @@ func New(ctx context.Context, config *Config, dc dns.Client, validator vless.Val
/* /*
if fb.Path != "" { if fb.Path != "" {
if r, err := regexp.Compile(fb.Path); err != nil { if r, err := regexp.Compile(fb.Path); err != nil {
return nil, errors.New("invalid path regexp").Base(err).AtError() return nil, errors.New("invalid path regexp").Base(err)
} else { } else {
handler.regexps[fb.Path] = r handler.regexps[fb.Path] = r
} }
@@ -274,13 +274,13 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
if h.decryption != nil { if h.decryption != nil {
var err error var err error
if connection, err = h.decryption.Handshake(connection, nil); err != nil { if connection, err = h.decryption.Handshake(connection, nil); err != nil {
return errors.New("ML-KEM-768 handshake failed").Base(err).AtInfo() return errors.New("ML-KEM-768 handshake failed").Base(err)
} }
} }
sessionPolicy := h.policyManager.ForLevel(0) sessionPolicy := h.policyManager.ForLevel(0)
if err := connection.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil { if err := connection.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
return errors.New("unable to set read deadline").Base(err).AtWarning() return errors.New("unable to set read deadline").Base(err)
} }
first := buf.FromBytes(make([]byte, buf.Size)) first := buf.FromBytes(make([]byte, buf.Size))
@@ -352,7 +352,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
} }
apfb := napfb[name] apfb := napfb[name]
if apfb == nil { if apfb == nil {
return errors.New(`failed to find the default "name" config`).AtWarning() return errors.New(`failed to find the default "name" config`)
} }
if apfb[alpn] == nil { if apfb[alpn] == nil {
@@ -360,7 +360,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
} }
pfb := apfb[alpn] pfb := apfb[alpn]
if pfb == nil { if pfb == nil {
return errors.New(`failed to find the default "alpn" config`).AtWarning() return errors.New(`failed to find the default "alpn" config`)
} }
path := "" path := ""
@@ -369,7 +369,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
if lines := bytes.Split(firstBytes, []byte{'\r', '\n'}); len(lines) > 1 { if lines := bytes.Split(firstBytes, []byte{'\r', '\n'}); len(lines) > 1 {
if s := bytes.Split(lines[0], []byte{' '}); len(s) == 3 { if s := bytes.Split(lines[0], []byte{' '}); len(s) == 3 {
if len(s[0]) < 8 && len(s[1]) > 0 && len(s[2]) == 8 { if len(s[0]) < 8 && len(s[1]) > 0 && len(s[2]) == 8 {
errors.New("realPath = " + string(s[1])).AtInfo().WriteToLog(sid) errors.New("realPath = " + string(s[1])).WriteToLog(sid)
for _, fb := range pfb { for _, fb := range pfb {
if fb.Path != "" && h.regexps[fb.Path].Match(s[1]) { if fb.Path != "" && h.regexps[fb.Path].Match(s[1]) {
path = fb.Path path = fb.Path
@@ -409,7 +409,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
} }
fb := pfb[path] fb := pfb[path]
if fb == nil { if fb == nil {
return errors.New(`failed to find the default "path" config`).AtWarning() return errors.New(`failed to find the default "path" config`)
} }
ctx, cancel := context.WithCancel(ctx) ctx, cancel := context.WithCancel(ctx)
@@ -425,7 +425,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
} }
return nil return nil
}); err != nil { }); err != nil {
return errors.New("failed to dial to " + fb.Dest).Base(err).AtWarning() return errors.New("failed to dial to " + fb.Dest).Base(err)
} }
defer conn.Close() defer conn.Close()
@@ -485,11 +485,11 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
pro.Write([]byte{byte(p1 >> 8), byte(p1), byte(p2 >> 8), byte(p2)}) pro.Write([]byte{byte(p1 >> 8), byte(p1), byte(p2 >> 8), byte(p2)})
} }
if err := serverWriter.WriteMultiBuffer(buf.MultiBuffer{pro}); err != nil { if err := serverWriter.WriteMultiBuffer(buf.MultiBuffer{pro}); err != nil {
return errors.New("failed to set PROXY protocol v", fb.Xver).Base(err).AtWarning() return errors.New("failed to set PROXY protocol v", fb.Xver).Base(err)
} }
} }
if err := buf.Copy(reader, serverWriter, buf.UpdateActivity(timer)); err != nil { if err := buf.Copy(reader, serverWriter, buf.UpdateActivity(timer)); err != nil {
return errors.New("failed to fallback request payload").Base(err).AtInfo() return errors.New("failed to fallback request payload").Base(err)
} }
return nil return nil
} }
@@ -499,7 +499,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
getResponse := func() error { getResponse := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
if err := buf.Copy(serverReader, writer, buf.UpdateActivity(timer)); err != nil { if err := buf.Copy(serverReader, writer, buf.UpdateActivity(timer)); err != nil {
return errors.New("failed to deliver response payload").Base(err).AtInfo() return errors.New("failed to deliver response payload").Base(err)
} }
return nil return nil
} }
@@ -507,7 +507,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
if err := task.Run(ctx, task.OnSuccess(postRequest, task.Close(serverWriter)), task.OnSuccess(getResponse, task.Close(writer))); err != nil { if err := task.Run(ctx, task.OnSuccess(postRequest, task.Close(serverWriter)), task.OnSuccess(getResponse, task.Close(writer))); err != nil {
common.Interrupt(serverReader) common.Interrupt(serverReader)
common.Interrupt(serverWriter) common.Interrupt(serverWriter)
return errors.New("fallback ends").Base(err).AtInfo() return errors.New("fallback ends").Base(err)
} }
return nil return nil
} }
@@ -519,7 +519,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
Status: log.AccessRejected, Status: log.AccessRejected,
Reason: err, Reason: err,
}) })
err = errors.New("invalid request from ", connection.RemoteAddr()).Base(err).AtInfo() err = errors.New("invalid request from ", connection.RemoteAddr()).Base(err)
} }
return err return err
} }
@@ -555,7 +555,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
inbound.CanSpliceCopy = 2 inbound.CanSpliceCopy = 2
switch request.Command { switch request.Command {
case protocol.RequestCommandUDP: case protocol.RequestCommandUDP:
return errors.New(requestAddons.Flow + " doesn't support UDP").AtWarning() return errors.New(requestAddons.Flow + " doesn't support UDP")
case protocol.RequestCommandMux, protocol.RequestCommandRvs: case protocol.RequestCommandMux, protocol.RequestCommandRvs:
inbound.CanSpliceCopy = 3 inbound.CanSpliceCopy = 3
fallthrough // we will break Mux connections that contain TCP requests fallthrough // we will break Mux connections that contain TCP requests
@@ -570,7 +570,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
p = uintptr(unsafe.Pointer(commonConn)) p = uintptr(unsafe.Pointer(commonConn))
} else if tlsConn, ok := iConn.(*tls.Conn); ok { } else if tlsConn, ok := iConn.(*tls.Conn); ok {
if tlsConn.ConnectionState().Version != gotls.VersionTLS13 { if tlsConn.ConnectionState().Version != gotls.VersionTLS13 {
return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, tlsConn.ConnectionState().Version).AtWarning() return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, tlsConn.ConnectionState().Version)
} }
t = reflect.TypeOf(tlsConn.Conn).Elem() t = reflect.TypeOf(tlsConn.Conn).Elem()
p = uintptr(unsafe.Pointer(tlsConn.Conn)) p = uintptr(unsafe.Pointer(tlsConn.Conn))
@@ -578,7 +578,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
t = reflect.TypeOf(realityConn.Conn).Elem() t = reflect.TypeOf(realityConn.Conn).Elem()
p = uintptr(unsafe.Pointer(realityConn.Conn)) p = uintptr(unsafe.Pointer(realityConn.Conn))
} else { } else {
return errors.New("XTLS only supports TLS and REALITY directly for now.").AtWarning() return errors.New("XTLS only supports TLS and REALITY directly for now.")
} }
i, _ := t.FieldByName("input") i, _ := t.FieldByName("input")
r, _ := t.FieldByName("rawInput") r, _ := t.FieldByName("rawInput")
@@ -586,15 +586,15 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
rawInput = (*bytes.Buffer)(unsafe.Pointer(p + r.Offset)) rawInput = (*bytes.Buffer)(unsafe.Pointer(p + r.Offset))
} }
} else { } else {
return errors.New("account " + account.ID.String() + " is not able to use the flow " + requestAddons.Flow).AtWarning() return errors.New("account " + account.ID.String() + " is not able to use the flow " + requestAddons.Flow)
} }
case "": case "":
inbound.CanSpliceCopy = 3 inbound.CanSpliceCopy = 3
if account.Flow == vless.XRV && (request.Command == protocol.RequestCommandTCP || isMuxAndNotXUDP(request, first)) { if account.Flow == vless.XRV && (request.Command == protocol.RequestCommandTCP || isMuxAndNotXUDP(request, first)) {
return errors.New("account " + account.ID.String() + " is rejected since the client flow is empty. Note that the pure TLS proxy has certain TLS in TLS characters.").AtWarning() return errors.New("account " + account.ID.String() + " is rejected since the client flow is empty. Note that the pure TLS proxy has certain TLS in TLS characters.")
} }
default: default:
return errors.New("unknown request flow " + requestAddons.Flow).AtWarning() return errors.New("unknown request flow " + requestAddons.Flow)
} }
if request.Command != protocol.RequestCommandMux { if request.Command != protocol.RequestCommandMux {
@@ -617,7 +617,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
bufferWriter := buf.NewBufferedWriter(buf.NewWriter(connection)) bufferWriter := buf.NewBufferedWriter(buf.NewWriter(connection))
if err := encoding.EncodeResponseHeader(bufferWriter, request, responseAddons); err != nil { if err := encoding.EncodeResponseHeader(bufferWriter, request, responseAddons); err != nil {
return errors.New("failed to encode response header").Base(err).AtWarning() return errors.New("failed to encode response header").Base(err)
} }
clientWriter := encoding.EncodeBodyAddons(bufferWriter, request, requestAddons, trafficState, false, ctx, connection, nil) clientWriter := encoding.EncodeBodyAddons(bufferWriter, request, requestAddons, trafficState, false, ctx, connection, nil)
bufferWriter.SetFlushNext() bufferWriter.SetFlushNext()
@@ -654,11 +654,11 @@ func (r *Reverse) Tag() string {
func (r *Reverse) NewMux(ctx context.Context, link *transport.Link, observer features.Feature) error { func (r *Reverse) NewMux(ctx context.Context, link *transport.Link, observer features.Feature) error {
muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{}) muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{})
if err != nil { if err != nil {
return errors.New("failed to create mux client worker").Base(err).AtWarning() return errors.New("failed to create mux client worker").Base(err)
} }
worker, err := reverse.NewPortalWorker(muxClient) worker, err := reverse.NewPortalWorker(muxClient)
if err != nil { if err != nil {
return errors.New("failed to create portal worker").Base(err).AtWarning() return errors.New("failed to create portal worker").Base(err)
} }
r.picker.AddWorker(worker) r.picker.AddWorker(worker)
if burstObs, ok := observer.(extension.BurstObservatory); ok { if burstObs, ok := observer.(extension.BurstObservatory); ok {
+18 -18
View File
@@ -73,7 +73,7 @@ func New(ctx context.Context, config *Config) (*Handler, error) {
} }
server, err := protocol.NewServerSpecFromPB(config.Vnext) server, err := protocol.NewServerSpecFromPB(config.Vnext)
if err != nil { if err != nil {
return nil, errors.New("failed to get server spec").Base(err).AtError() return nil, errors.New("failed to get server spec").Base(err)
} }
v := core.MustFromContext(ctx) v := core.MustFromContext(ctx)
@@ -93,7 +93,7 @@ func New(ctx context.Context, config *Config) (*Handler, error) {
} }
handler.encryption = &encryption.ClientInstance{} handler.encryption = &encryption.ClientInstance{}
if err := handler.encryption.Init(nfsPKeysBytes, a.XorMode, a.Seconds, a.Padding); err != nil { if err := handler.encryption.Init(nfsPKeysBytes, a.XorMode, a.Seconds, a.Padding); err != nil {
return nil, errors.New("failed to use encryption").Base(err).AtError() return nil, errors.New("failed to use encryption").Base(err)
} }
} }
@@ -106,7 +106,7 @@ func New(ctx context.Context, config *Config) (*Handler, error) {
if sc := a.Reverse.Sniffing; sc != nil && sc.Enabled { if sc := a.Reverse.Sniffing; sc != nil && sc.Enabled {
request, err := proxymanConfig.BuildSniffingRequest(sc) request, err := proxymanConfig.BuildSniffingRequest(sc)
if err != nil { if err != nil {
return nil, errors.New("failed to build reverse sniffing request").Base(err).AtError() return nil, errors.New("failed to build reverse sniffing request").Base(err)
} }
rvsCtx = session.ContextWithContent(rvsCtx, &session.Content{ rvsCtx = session.ContextWithContent(rvsCtx, &session.Content{
SniffingRequest: request, SniffingRequest: request,
@@ -149,7 +149,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
outbounds := session.OutboundsFromContext(ctx) outbounds := session.OutboundsFromContext(ctx)
ob := outbounds[len(outbounds)-1] ob := outbounds[len(outbounds)-1]
if !ob.Target.IsValid() && ob.Target.Address.String() != "v1.rvs.cool" { if !ob.Target.IsValid() && ob.Target.Address.String() != "v1.rvs.cool" {
return errors.New("target not specified").AtError() return errors.New("target not specified")
} }
ob.Name = "vless" ob.Name = "vless"
@@ -178,7 +178,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
for { for {
connTime := <-h.preConns connTime := <-h.preConns
if connTime == nil { if connTime == nil {
return errors.New("closed handler").AtWarning() return errors.New("closed handler")
} }
if time.Now().Before(connTime.Expire) { if time.Now().Before(connTime.Expire) {
conn = connTime.Conn conn = connTime.Conn
@@ -197,7 +197,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
} }
return nil return nil
}); err != nil { }); err != nil {
return errors.New("failed to find an available destination").Base(err).AtWarning() return errors.New("failed to find an available destination").Base(err)
} }
} }
defer conn.Close() defer conn.Close()
@@ -209,7 +209,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
if h.encryption != nil { if h.encryption != nil {
var err error var err error
if conn, err = h.encryption.Handshake(conn); err != nil { if conn, err = h.encryption.Handshake(conn); err != nil {
return errors.New("ML-KEM-768 handshake failed").Base(err).AtInfo() return errors.New("ML-KEM-768 handshake failed").Base(err)
} }
} }
@@ -223,7 +223,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
command = protocol.RequestCommandMux command = protocol.RequestCommandMux
case "v1.rvs.cool": case "v1.rvs.cool":
if target.Network != net.Network_Unknown { if target.Network != net.Network_Unknown {
return errors.New("nice try baby").AtError() return errors.New("nice try baby")
} }
command = protocol.RequestCommandRvs command = protocol.RequestCommandRvs
} }
@@ -256,7 +256,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
switch request.Command { switch request.Command {
case protocol.RequestCommandUDP: case protocol.RequestCommandUDP:
if !allowUDP443 && request.Port == 443 { if !allowUDP443 && request.Port == 443 {
return errors.New("XTLS rejected UDP/443 traffic").AtInfo() return errors.New("XTLS rejected UDP/443 traffic")
} }
case protocol.RequestCommandMux: case protocol.RequestCommandMux:
fallthrough // let server break Mux connections that contain TCP requests fallthrough // let server break Mux connections that contain TCP requests
@@ -279,7 +279,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
t = reflect.TypeOf(realityConn.Conn).Elem() t = reflect.TypeOf(realityConn.Conn).Elem()
p = uintptr(unsafe.Pointer(realityConn.Conn)) p = uintptr(unsafe.Pointer(realityConn.Conn))
} else { } else {
return errors.New("XTLS only supports TLS and REALITY directly for now.").AtWarning() return errors.New("XTLS only supports TLS and REALITY directly for now.")
} }
i, _ := t.FieldByName("input") i, _ := t.FieldByName("input")
r, _ := t.FieldByName("rawInput") r, _ := t.FieldByName("rawInput")
@@ -321,7 +321,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
bufferWriter := buf.NewBufferedWriter(buf.NewWriter(conn)) bufferWriter := buf.NewBufferedWriter(buf.NewWriter(conn))
if err := encoding.EncodeRequestHeader(bufferWriter, request, requestAddons); err != nil { if err := encoding.EncodeRequestHeader(bufferWriter, request, requestAddons); err != nil {
return errors.New("failed to encode request header").Base(err).AtWarning() return errors.New("failed to encode request header").Base(err)
} }
// default: serverWriter := bufferWriter // default: serverWriter := bufferWriter
@@ -350,23 +350,23 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
} }
// Flush; bufferWriter.WriteMultiBuffer now is bufferWriter.writer.WriteMultiBuffer // Flush; bufferWriter.WriteMultiBuffer now is bufferWriter.writer.WriteMultiBuffer
if err := bufferWriter.SetBuffered(false); err != nil { if err := bufferWriter.SetBuffered(false); err != nil {
return errors.New("failed to write A request payload").Base(err).AtWarning() return errors.New("failed to write A request payload").Base(err)
} }
if requestAddons.Flow == vless.XRV { if requestAddons.Flow == vless.XRV {
if tlsConn, ok := iConn.(*tls.Conn); ok { if tlsConn, ok := iConn.(*tls.Conn); ok {
if tlsConn.ConnectionState().Version != gotls.VersionTLS13 { if tlsConn.ConnectionState().Version != gotls.VersionTLS13 {
return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, tlsConn.ConnectionState().Version).AtWarning() return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, tlsConn.ConnectionState().Version)
} }
} else if utlsConn, ok := iConn.(*tls.UConn); ok { } else if utlsConn, ok := iConn.(*tls.UConn); ok {
if utlsConn.ConnectionState().Version != utls.VersionTLS13 { if utlsConn.ConnectionState().Version != utls.VersionTLS13 {
return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, utlsConn.ConnectionState().Version).AtWarning() return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, utlsConn.ConnectionState().Version)
} }
} }
} }
err := buf.Copy(clientReader, serverWriter, buf.UpdateActivity(timer)) err := buf.Copy(clientReader, serverWriter, buf.UpdateActivity(timer))
if err != nil { if err != nil {
return errors.New("failed to transfer request payload").Base(err).AtInfo() return errors.New("failed to transfer request payload").Base(err)
} }
// Indicates the end of request payload. // Indicates the end of request payload.
@@ -381,7 +381,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
responseAddons, err := encoding.DecodeResponseHeader(conn, request) responseAddons, err := encoding.DecodeResponseHeader(conn, request)
if err != nil { if err != nil {
return errors.New("failed to decode response header").Base(err).AtInfo() return errors.New("failed to decode response header").Base(err)
} }
// default: serverReader := buf.NewReader(conn) // default: serverReader := buf.NewReader(conn)
@@ -405,7 +405,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
} }
if err != nil { if err != nil {
return errors.New("failed to transfer response payload").Base(err).AtInfo() return errors.New("failed to transfer response payload").Base(err)
} }
return nil return nil
@@ -416,7 +416,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
} }
if err := task.Run(ctx, postRequest, task.OnSuccess(getResponse, task.Close(clientWriter))); err != nil { if err := task.Run(ctx, postRequest, task.OnSuccess(getResponse, task.Close(clientWriter))); err != nil {
return errors.New("connection ends").Base(err).AtInfo() return errors.New("connection ends").Base(err)
} }
return nil return nil
+1 -1
View File
@@ -49,7 +49,7 @@ func (a *MemoryAccount) ToProto() proto.Message {
func (a *Account) AsAccount() (protocol.Account, error) { func (a *Account) AsAccount() (protocol.Account, error) {
id, err := uuid.ParseString(a.Id) id, err := uuid.ParseString(a.Id)
if err != nil { if err != nil {
return nil, errors.New("failed to parse ID").Base(err).AtError() return nil, errors.New("failed to parse ID").Base(err)
} }
protoID := protocol.NewID(id) protoID := protocol.NewID(id)
var AuthenticatedLength, NoTerminationSignal bool var AuthenticatedLength, NoTerminationSignal bool
+1 -1
View File
@@ -209,7 +209,7 @@ func (c *ClientSession) DecodeResponseHeader(reader io.Reader) (*protocol.Respon
defer buffer.Release() defer buffer.Release()
if _, err := buffer.ReadFullFrom(c.responseReader, 4); err != nil { if _, err := buffer.ReadFullFrom(c.responseReader, 4); err != nil {
return nil, errors.New("failed to read response header").Base(err).AtWarning() return nil, errors.New("failed to read response header").Base(err)
} }
if buffer.Byte(0) != c.responseHeader { if buffer.Byte(0) != c.responseHeader {
+2 -2
View File
@@ -227,7 +227,7 @@ func transferResponse(timer signal.ActivityUpdater, session *encoding.ServerSess
func (h *Handler) Process(ctx context.Context, network net.Network, connection stat.Connection, dispatcher routing.Dispatcher) error { func (h *Handler) Process(ctx context.Context, network net.Network, connection stat.Connection, dispatcher routing.Dispatcher) error {
sessionPolicy := h.policyManager.ForLevel(0) sessionPolicy := h.policyManager.ForLevel(0)
if err := connection.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil { if err := connection.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
return errors.New("unable to set read deadline").Base(err).AtWarning() return errors.New("unable to set read deadline").Base(err)
} }
iConn := stat.TryUnwrapStatsConn(connection) iConn := stat.TryUnwrapStatsConn(connection)
@@ -247,7 +247,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s
Status: log.AccessRejected, Status: log.AccessRejected,
Reason: err, Reason: err,
}) })
err = errors.New("invalid request from ", connection.RemoteAddr()).Base(err).AtInfo() err = errors.New("invalid request from ", connection.RemoteAddr()).Base(err)
} }
return err return err
} }
+3 -3
View File
@@ -60,7 +60,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
outbounds := session.OutboundsFromContext(ctx) outbounds := session.OutboundsFromContext(ctx)
ob := outbounds[len(outbounds)-1] ob := outbounds[len(outbounds)-1]
if !ob.Target.IsValid() { if !ob.Target.IsValid() {
return errors.New("target not specified").AtError() return errors.New("target not specified")
} }
ob.Name = "vmess" ob.Name = "vmess"
ob.CanSpliceCopy = 3 ob.CanSpliceCopy = 3
@@ -78,7 +78,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
return nil return nil
}) })
if err != nil { if err != nil {
return errors.New("failed to find an available destination").Base(err).AtWarning() return errors.New("failed to find an available destination").Base(err)
} }
defer conn.Close() defer conn.Close()
@@ -154,7 +154,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
writer := buf.NewBufferedWriter(buf.NewWriter(conn)) writer := buf.NewBufferedWriter(buf.NewWriter(conn))
if err := session.EncodeRequestHeader(request, writer); err != nil { if err := session.EncodeRequestHeader(request, writer); err != nil {
return errors.New("failed to encode request").Base(err).AtWarning() return errors.New("failed to encode request").Base(err)
} }
bodyWriter, err := session.EncodeRequestBody(request, writer) bodyWriter, err := session.EncodeRequestBody(request, writer)
+138 -113
View File
@@ -3,6 +3,7 @@ package wireguard
import ( import (
"context" "context"
"fmt" "fmt"
gonet "net"
"net/netip" "net/netip"
"reflect" "reflect"
"strings" "strings"
@@ -27,10 +28,14 @@ import (
"github.com/xtls/xray-core/features/stats" "github.com/xtls/xray-core/features/stats"
"github.com/xtls/xray-core/transport" "github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet" "github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/finalmask"
"golang.zx2c4.com/wireguard/device" "golang.zx2c4.com/wireguard/device"
) )
type entry struct {
got []net.IP
time time.Time
}
type Handler struct { type Handler struct {
conf *DeviceConfig conf *DeviceConfig
policyManager policy.Manager policyManager policy.Manager
@@ -44,6 +49,11 @@ type Handler struct {
tnet *Net tnet *Net
dev *device.Device dev *device.Device
mu sync.Mutex mu sync.Mutex
// TODO: cache cleanup loop
local bool
cache map[string]entry
cacheMu sync.Mutex
} }
func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) { func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
@@ -99,10 +109,15 @@ func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
return nil, err return nil, err
} }
local := false
dns := conf.DNS dns := conf.DNS
if len(dns) == 0 { if len(dns) == 0 {
dns = []string{"1.1.1.1", "1.0.0.1", "2606:4700:4700::1111", "2606:4700:4700::1001"} dns = []string{"1.1.1.1", "1.0.0.1", "2606:4700:4700::1111", "2606:4700:4700::1001"}
} }
if len(dns) == 1 && dns[0] == "local" {
local = true
dns = nil
}
dnses := make([]netip.Addr, 0, len(dns)) dnses := make([]netip.Addr, 0, len(dns))
for _, dns := range dns { for _, dns := range dns {
dnses = append(dnses, netip.MustParseAddr(dns)) dnses = append(dnses, netip.MustParseAddr(dns))
@@ -136,6 +151,9 @@ func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
tun: tun, tun: tun,
tnet: tnet, tnet: tnet,
local: local,
cache: make(map[string]entry),
}, nil }, nil
} }
@@ -154,6 +172,22 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
return err return err
} }
var addr netip.Addr
if ob.Target.Address.Family().IsDomain() {
ip, err := h.resolveRemote(ob.Target.Address.String())
if err != nil {
return errors.New("failed to resolve domain").Base(err)
}
addr, _ = netip.AddrFromSlice(ip)
} else {
addr, _ = netip.AddrFromSlice(ob.Target.Address.IP())
}
addrPort := netip.AddrPortFrom(addr, ob.Target.Port.Value())
if !addrPort.IsValid() {
return errors.New("invalid target ", ob.Target)
}
var newCtx context.Context var newCtx context.Context
var newCancel context.CancelFunc var newCancel context.CancelFunc
if session.TimeoutOnlyFromContext(ctx) { if session.TimeoutOnlyFromContext(ctx) {
@@ -182,10 +216,10 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
var err error var err error
if sessionPolicy.Timeouts.Handshake != 0 { if sessionPolicy.Timeouts.Handshake != 0 {
timeoutCtx, timeoutCancel := context.WithTimeout(ctx, sessionPolicy.Timeouts.Handshake) timeoutCtx, timeoutCancel := context.WithTimeout(ctx, sessionPolicy.Timeouts.Handshake)
conn, err = h.tnet.DialContext(timeoutCtx, "tcp", ob.Target.NetAddr()) conn, err = h.tnet.DialContextTCPAddrPort(timeoutCtx, addrPort)
timeoutCancel() timeoutCancel()
} else { } else {
conn, err = h.tnet.Dial("tcp", ob.Target.NetAddr()) conn, err = h.tnet.DialContextTCPAddrPort(ctx, addrPort)
} }
if err != nil { if err != nil {
return errors.New("failed to create TCP connection").Base(err) return errors.New("failed to create TCP connection").Base(err)
@@ -194,14 +228,15 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
reader = buf.NewReader(conn) reader = buf.NewReader(conn)
writer = buf.NewWriter(conn) writer = buf.NewWriter(conn)
case net.Network_UDP: case net.Network_UDP:
conn, err := h.tnet.Dial("udp", ob.Target.NetAddr()) conn, err := h.tnet.DialUDPAddrPort(netip.AddrPort{}, addrPort)
if err != nil { if err != nil {
return errors.New("failed to create UDP connection").Base(err) return errors.New("failed to create UDP connection").Base(err)
} }
defer conn.Close() defer conn.Close()
c := &udpConnClient{ c := &udpConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn, PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
dest: conn.RemoteAddr().(*net.UDPAddr), resolveFunc: h.resolveRemote,
dest: gonet.UDPAddrFromAddrPort(addrPort),
} }
reader = c reader = c
writer = c writer = c
@@ -258,26 +293,26 @@ func (h *Handler) init(ctx context.Context) error {
if err != nil { if err != nil {
return nil, err return nil, err
} }
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
if err != nil {
return nil, err
}
var pktConn net.PacketConn var pktConn net.PacketConn
if h.streamSettings.FinalMask != nil { switch c := conn.(type) {
conn, err := h.streamSettings.FinalMask.DialUDP(ctx, dest) case *internet.PacketConnWrapper:
pktConn = c.PacketConn
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
default:
panic(reflect.TypeOf(c))
}
if h.streamSettings.UdpmaskManager != nil {
newConn, err := h.streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn)
if err != nil { if err != nil {
return nil, errors.New("failed to dial to dest").Base(err) pktConn.Close()
} return nil, errors.New("mask err").Base(err)
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
} else {
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
if err != nil {
return nil, errors.New("failed to dial to dest").Base(err)
}
switch c := conn.(type) {
case *internet.PacketConnWrapper:
pktConn = c.PacketConn
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
default:
panic(reflect.TypeOf(c))
} }
pktConn = newConn
} }
if h.uplinkCounter != nil || h.downlinkCounter != nil { if h.uplinkCounter != nil || h.downlinkCounter != nil {
pktConn = &PacketCounterConnection{ pktConn = &PacketCounterConnection{
@@ -336,48 +371,87 @@ func (h *Handler) init(ctx context.Context) error {
} }
func (h *Handler) resolveLocal(host string) (net.IP, error) { func (h *Handler) resolveLocal(host string) (net.IP, error) {
ips, _, err := h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true}) return h.resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, uint32, error) {
return h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
})
}
func (h *Handler) resolveRemote(host string) (net.IP, error) {
return h.resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, uint32, error) {
if h.local {
return h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
}
return h.tnet.LookupHost(host)
})
}
func (h *Handler) resolveDomain(host string, strategy DeviceConfig_DomainStrategy, lookupIP func(host string) ([]net.IP, uint32, error)) (net.IP, error) {
if ip := net.ParseIP(host); ip != nil {
return ip, nil
}
h.cacheMu.Lock()
if entry, ok := h.cache[host]; ok {
if time.Now().Before(entry.time) {
h.cacheMu.Unlock()
return entry.got[dice.Roll(len(entry.got))], nil
}
delete(h.cache, host)
}
h.cacheMu.Unlock()
ips, ttl, err := lookupIP(host)
if err != nil { if err != nil {
return nil, err return nil, err
} }
got := ips if len(ips) == 0 {
if h.streamSettings.SocketSettings != nil { return nil, dns.ErrEmptyResponse
var got4, got6 []net.IP }
for _, ip := range ips { var got4, got6 []net.IP
if ip.To4() != nil { for _, ip := range ips {
got4 = append(got4, ip) if ip.To4() != nil {
} else { got4 = append(got4, ip)
got6 = append(got6, ip) } else {
} got6 = append(got6, ip)
}
switch h.streamSettings.SocketSettings.DomainStrategy {
case internet.DomainStrategy_AS_IS, internet.DomainStrategy_USE_IP, internet.DomainStrategy_FORCE_IP:
got = ips
case internet.DomainStrategy_USE_IP4, internet.DomainStrategy_FORCE_IP4:
got = got4
case internet.DomainStrategy_USE_IP6, internet.DomainStrategy_FORCE_IP6:
got = got6
case internet.DomainStrategy_USE_IP46, internet.DomainStrategy_FORCE_IP46:
got = got4
if len(got) == 0 {
got = got6
}
case internet.DomainStrategy_USE_IP64, internet.DomainStrategy_FORCE_IP64:
got = got6
if len(got) == 0 {
got = got4
}
}
if len(got) == 0 {
return nil, dns.ErrEmptyResponse
} }
} }
var got []net.IP
switch strategy {
case DeviceConfig_FORCE_IP:
got = ips
return ips[dice.Roll(len(ips))], nil
case DeviceConfig_FORCE_IP4:
got = got4
case DeviceConfig_FORCE_IP6:
got = got6
case DeviceConfig_FORCE_IP46:
got = got4
if len(got) == 0 {
got = got6
}
case DeviceConfig_FORCE_IP64:
got = got6
if len(got) == 0 {
got = got4
}
default:
panic(strategy)
}
if len(got) == 0 {
return nil, dns.ErrEmptyResponse
}
entry := entry{
got: got,
time: time.Now().Add(time.Duration(ttl) * time.Second),
}
h.cacheMu.Lock()
h.cache[host] = entry
h.cacheMu.Unlock()
return got[dice.Roll(len(got))], nil return got[dice.Roll(len(got))], nil
} }
type udpConnClient struct { type udpConnClient struct {
net.PacketConn net.PacketConn
dest *net.UDPAddr resolveFunc func(host string) (net.IP, error)
dest *net.UDPAddr
} }
func (c *udpConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) { func (c *udpConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
@@ -404,8 +478,15 @@ func (c *udpConnClient) WriteMultiBuffer(mb buf.MultiBuffer) error {
dst := c.dest dst := c.dest
if b.UDP != nil { if b.UDP != nil {
if b.UDP.Address.Family().IsDomain() { if b.UDP.Address.Family().IsDomain() {
if b.UDP.Port != net.Port(dst.Port) { ip, err := c.resolveFunc(b.UDP.Address.String())
dst = &net.UDPAddr{IP: dst.IP, Port: int(b.UDP.Port)} if err != nil {
errors.LogErrorInner(context.Background(), err, "drop packet to ", b.UDP, " with size ", len(b.Bytes()))
b.Release()
continue
}
dst = &net.UDPAddr{
IP: ip,
Port: int(b.UDP.Port),
} }
} else { } else {
dst = b.UDP.RawNetAddr().(*net.UDPAddr) dst = b.UDP.RawNetAddr().(*net.UDPAddr)
@@ -442,59 +523,3 @@ func (c *PacketCounterConnection) WriteTo(p []byte, addr net.Addr) (n int, err e
} }
return return
} }
type entry struct {
saddr []string
deadline time.Time
}
type cache struct {
running bool
m map[string]entry
mu sync.Mutex
}
func (c *cache) run() {
if c.running {
return
}
c.running = true
c.m = make(map[string]entry)
go c.gc()
}
func (c *cache) gc() {
ticker := time.NewTicker(time.Minute)
for {
now := <-ticker.C
c.mu.Lock()
for key, entry := range c.m {
if now.After(entry.deadline) {
delete(c.m, key)
}
}
c.mu.Unlock()
}
}
func (c *cache) LookupHost(host string) []string {
c.mu.Lock()
defer c.mu.Unlock()
c.run()
if entry, ok := c.m[host]; ok {
if time.Now().Before(entry.deadline) {
return entry.saddr
}
delete(c.m, host)
}
return nil
}
func (c *cache) Cache(host string, saddr []string, ttl uint32) {
c.mu.Lock()
defer c.mu.Unlock()
c.m[host] = entry{
saddr: saddr,
deadline: time.Now().Add(time.Second * time.Duration(ttl)),
}
}
+102 -26
View File
@@ -22,6 +22,61 @@ const (
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
) )
type DeviceConfig_DomainStrategy int32
const (
DeviceConfig_FORCE_IP DeviceConfig_DomainStrategy = 0
DeviceConfig_FORCE_IP4 DeviceConfig_DomainStrategy = 1
DeviceConfig_FORCE_IP6 DeviceConfig_DomainStrategy = 2
DeviceConfig_FORCE_IP46 DeviceConfig_DomainStrategy = 3
DeviceConfig_FORCE_IP64 DeviceConfig_DomainStrategy = 4
)
// Enum value maps for DeviceConfig_DomainStrategy.
var (
DeviceConfig_DomainStrategy_name = map[int32]string{
0: "FORCE_IP",
1: "FORCE_IP4",
2: "FORCE_IP6",
3: "FORCE_IP46",
4: "FORCE_IP64",
}
DeviceConfig_DomainStrategy_value = map[string]int32{
"FORCE_IP": 0,
"FORCE_IP4": 1,
"FORCE_IP6": 2,
"FORCE_IP46": 3,
"FORCE_IP64": 4,
}
)
func (x DeviceConfig_DomainStrategy) Enum() *DeviceConfig_DomainStrategy {
p := new(DeviceConfig_DomainStrategy)
*p = x
return p
}
func (x DeviceConfig_DomainStrategy) String() string {
return protoimpl.X.EnumStringOf(x.Descriptor(), protoreflect.EnumNumber(x))
}
func (DeviceConfig_DomainStrategy) Descriptor() protoreflect.EnumDescriptor {
return file_proxy_wireguard_config_proto_enumTypes[0].Descriptor()
}
func (DeviceConfig_DomainStrategy) Type() protoreflect.EnumType {
return &file_proxy_wireguard_config_proto_enumTypes[0]
}
func (x DeviceConfig_DomainStrategy) Number() protoreflect.EnumNumber {
return protoreflect.EnumNumber(x)
}
// Deprecated: Use DeviceConfig_DomainStrategy.Descriptor instead.
func (DeviceConfig_DomainStrategy) EnumDescriptor() ([]byte, []int) {
return file_proxy_wireguard_config_proto_rawDescGZIP(), []int{1, 0}
}
type PeerConfig struct { type PeerConfig struct {
state protoimpl.MessageState `protogen:"open.v1"` state protoimpl.MessageState `protogen:"open.v1"`
PublicKey string `protobuf:"bytes,1,opt,name=public_key,json=publicKey,proto3" json:"public_key,omitempty"` PublicKey string `protobuf:"bytes,1,opt,name=public_key,json=publicKey,proto3" json:"public_key,omitempty"`
@@ -99,18 +154,19 @@ func (x *PeerConfig) GetAllowedIps() []string {
} }
type DeviceConfig struct { type DeviceConfig struct {
state protoimpl.MessageState `protogen:"open.v1"` state protoimpl.MessageState `protogen:"open.v1"`
SecretKey string `protobuf:"bytes,1,opt,name=secret_key,json=secretKey,proto3" json:"secret_key,omitempty"` SecretKey string `protobuf:"bytes,1,opt,name=secret_key,json=secretKey,proto3" json:"secret_key,omitempty"`
Endpoint []string `protobuf:"bytes,2,rep,name=endpoint,proto3" json:"endpoint,omitempty"` Endpoint []string `protobuf:"bytes,2,rep,name=endpoint,proto3" json:"endpoint,omitempty"`
Peers []*PeerConfig `protobuf:"bytes,3,rep,name=peers,proto3" json:"peers,omitempty"` Peers []*PeerConfig `protobuf:"bytes,3,rep,name=peers,proto3" json:"peers,omitempty"`
Users []*protocol.User `protobuf:"bytes,5,rep,name=users,proto3" json:"users,omitempty"` Users []*protocol.User `protobuf:"bytes,5,rep,name=users,proto3" json:"users,omitempty"`
Mtu int32 `protobuf:"varint,4,opt,name=mtu,proto3" json:"mtu,omitempty"` Mtu int32 `protobuf:"varint,4,opt,name=mtu,proto3" json:"mtu,omitempty"`
Reserved []byte `protobuf:"bytes,6,opt,name=reserved,proto3" json:"reserved,omitempty"` Reserved []byte `protobuf:"bytes,6,opt,name=reserved,proto3" json:"reserved,omitempty"`
IsClient bool `protobuf:"varint,8,opt,name=is_client,json=isClient,proto3" json:"is_client,omitempty"` DomainStrategy DeviceConfig_DomainStrategy `protobuf:"varint,7,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.proxy.wireguard.DeviceConfig_DomainStrategy" json:"domain_strategy,omitempty"`
NoKernelTun bool `protobuf:"varint,9,opt,name=no_kernel_tun,json=noKernelTun,proto3" json:"no_kernel_tun,omitempty"` IsClient bool `protobuf:"varint,8,opt,name=is_client,json=isClient,proto3" json:"is_client,omitempty"`
DNS []string `protobuf:"bytes,10,rep,name=DNS,proto3" json:"DNS,omitempty"` NoKernelTun bool `protobuf:"varint,9,opt,name=no_kernel_tun,json=noKernelTun,proto3" json:"no_kernel_tun,omitempty"`
unknownFields protoimpl.UnknownFields DNS []string `protobuf:"bytes,10,rep,name=DNS,proto3" json:"DNS,omitempty"`
sizeCache protoimpl.SizeCache unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
} }
func (x *DeviceConfig) Reset() { func (x *DeviceConfig) Reset() {
@@ -185,6 +241,13 @@ func (x *DeviceConfig) GetReserved() []byte {
return nil return nil
} }
func (x *DeviceConfig) GetDomainStrategy() DeviceConfig_DomainStrategy {
if x != nil {
return x.DomainStrategy
}
return DeviceConfig_FORCE_IP
}
func (x *DeviceConfig) GetIsClient() bool { func (x *DeviceConfig) GetIsClient() bool {
if x != nil { if x != nil {
return x.IsClient return x.IsClient
@@ -220,7 +283,7 @@ const file_proxy_wireguard_config_proto_rawDesc = "" +
"\n" + "\n" +
"keep_alive\x18\x04 \x01(\tR\tkeepAlive\x12\x1f\n" + "keep_alive\x18\x04 \x01(\tR\tkeepAlive\x12\x1f\n" +
"\vallowed_ips\x18\x05 \x03(\tR\n" + "\vallowed_ips\x18\x05 \x03(\tR\n" +
"allowedIps\"\xb4\x02\n" + "allowedIps\"\xee\x03\n" +
"\fDeviceConfig\x12\x1d\n" + "\fDeviceConfig\x12\x1d\n" +
"\n" + "\n" +
"secret_key\x18\x01 \x01(\tR\tsecretKey\x12\x1a\n" + "secret_key\x18\x01 \x01(\tR\tsecretKey\x12\x1a\n" +
@@ -228,11 +291,20 @@ const file_proxy_wireguard_config_proto_rawDesc = "" +
"\x05peers\x18\x03 \x03(\v2 .xray.proxy.wireguard.PeerConfigR\x05peers\x120\n" + "\x05peers\x18\x03 \x03(\v2 .xray.proxy.wireguard.PeerConfigR\x05peers\x120\n" +
"\x05users\x18\x05 \x03(\v2\x1a.xray.common.protocol.UserR\x05users\x12\x10\n" + "\x05users\x18\x05 \x03(\v2\x1a.xray.common.protocol.UserR\x05users\x12\x10\n" +
"\x03mtu\x18\x04 \x01(\x05R\x03mtu\x12\x1a\n" + "\x03mtu\x18\x04 \x01(\x05R\x03mtu\x12\x1a\n" +
"\breserved\x18\x06 \x01(\fR\breserved\x12\x1b\n" + "\breserved\x18\x06 \x01(\fR\breserved\x12Z\n" +
"\x0fdomain_strategy\x18\a \x01(\x0e21.xray.proxy.wireguard.DeviceConfig.DomainStrategyR\x0edomainStrategy\x12\x1b\n" +
"\tis_client\x18\b \x01(\bR\bisClient\x12\"\n" + "\tis_client\x18\b \x01(\bR\bisClient\x12\"\n" +
"\rno_kernel_tun\x18\t \x01(\bR\vnoKernelTun\x12\x10\n" + "\rno_kernel_tun\x18\t \x01(\bR\vnoKernelTun\x12\x10\n" +
"\x03DNS\x18\n" + "\x03DNS\x18\n" +
" \x03(\tR\x03DNSB^\n" + " \x03(\tR\x03DNS\"\\\n" +
"\x0eDomainStrategy\x12\f\n" +
"\bFORCE_IP\x10\x00\x12\r\n" +
"\tFORCE_IP4\x10\x01\x12\r\n" +
"\tFORCE_IP6\x10\x02\x12\x0e\n" +
"\n" +
"FORCE_IP46\x10\x03\x12\x0e\n" +
"\n" +
"FORCE_IP64\x10\x04B^\n" +
"\x18com.xray.proxy.wireguardP\x01Z)github.com/xtls/xray-core/proxy/wireguard\xaa\x02\x14Xray.Proxy.WireGuardb\x06proto3" "\x18com.xray.proxy.wireguardP\x01Z)github.com/xtls/xray-core/proxy/wireguard\xaa\x02\x14Xray.Proxy.WireGuardb\x06proto3"
var ( var (
@@ -247,20 +319,23 @@ func file_proxy_wireguard_config_proto_rawDescGZIP() []byte {
return file_proxy_wireguard_config_proto_rawDescData return file_proxy_wireguard_config_proto_rawDescData
} }
var file_proxy_wireguard_config_proto_enumTypes = make([]protoimpl.EnumInfo, 1)
var file_proxy_wireguard_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2) var file_proxy_wireguard_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2)
var file_proxy_wireguard_config_proto_goTypes = []any{ var file_proxy_wireguard_config_proto_goTypes = []any{
(*PeerConfig)(nil), // 0: xray.proxy.wireguard.PeerConfig (DeviceConfig_DomainStrategy)(0), // 0: xray.proxy.wireguard.DeviceConfig.DomainStrategy
(*DeviceConfig)(nil), // 1: xray.proxy.wireguard.DeviceConfig (*PeerConfig)(nil), // 1: xray.proxy.wireguard.PeerConfig
(*protocol.User)(nil), // 2: xray.common.protocol.User (*DeviceConfig)(nil), // 2: xray.proxy.wireguard.DeviceConfig
(*protocol.User)(nil), // 3: xray.common.protocol.User
} }
var file_proxy_wireguard_config_proto_depIdxs = []int32{ var file_proxy_wireguard_config_proto_depIdxs = []int32{
0, // 0: xray.proxy.wireguard.DeviceConfig.peers:type_name -> xray.proxy.wireguard.PeerConfig 1, // 0: xray.proxy.wireguard.DeviceConfig.peers:type_name -> xray.proxy.wireguard.PeerConfig
2, // 1: xray.proxy.wireguard.DeviceConfig.users:type_name -> xray.common.protocol.User 3, // 1: xray.proxy.wireguard.DeviceConfig.users:type_name -> xray.common.protocol.User
2, // [2:2] is the sub-list for method output_type 0, // 2: xray.proxy.wireguard.DeviceConfig.domain_strategy:type_name -> xray.proxy.wireguard.DeviceConfig.DomainStrategy
2, // [2:2] is the sub-list for method input_type 3, // [3:3] is the sub-list for method output_type
2, // [2:2] is the sub-list for extension type_name 3, // [3:3] is the sub-list for method input_type
2, // [2:2] is the sub-list for extension extendee 3, // [3:3] is the sub-list for extension type_name
0, // [0:2] is the sub-list for field type_name 3, // [3:3] is the sub-list for extension extendee
0, // [0:3] is the sub-list for field type_name
} }
func init() { file_proxy_wireguard_config_proto_init() } func init() { file_proxy_wireguard_config_proto_init() }
@@ -273,13 +348,14 @@ func file_proxy_wireguard_config_proto_init() {
File: protoimpl.DescBuilder{ File: protoimpl.DescBuilder{
GoPackagePath: reflect.TypeOf(x{}).PkgPath(), GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_wireguard_config_proto_rawDesc), len(file_proxy_wireguard_config_proto_rawDesc)), RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_wireguard_config_proto_rawDesc), len(file_proxy_wireguard_config_proto_rawDesc)),
NumEnums: 0, NumEnums: 1,
NumMessages: 2, NumMessages: 2,
NumExtensions: 0, NumExtensions: 0,
NumServices: 0, NumServices: 0,
}, },
GoTypes: file_proxy_wireguard_config_proto_goTypes, GoTypes: file_proxy_wireguard_config_proto_goTypes,
DependencyIndexes: file_proxy_wireguard_config_proto_depIdxs, DependencyIndexes: file_proxy_wireguard_config_proto_depIdxs,
EnumInfos: file_proxy_wireguard_config_proto_enumTypes,
MessageInfos: file_proxy_wireguard_config_proto_msgTypes, MessageInfos: file_proxy_wireguard_config_proto_msgTypes,
}.Build() }.Build()
File_proxy_wireguard_config_proto = out.File File_proxy_wireguard_config_proto = out.File
+8
View File
@@ -17,6 +17,13 @@ message PeerConfig {
} }
message DeviceConfig { message DeviceConfig {
enum DomainStrategy {
FORCE_IP = 0;
FORCE_IP4 = 1;
FORCE_IP6 = 2;
FORCE_IP46 = 3;
FORCE_IP64 = 4;
}
string secret_key = 1; string secret_key = 1;
repeated string endpoint = 2; repeated string endpoint = 2;
repeated PeerConfig peers = 3; repeated PeerConfig peers = 3;
@@ -24,6 +31,7 @@ message DeviceConfig {
int32 mtu = 4; int32 mtu = 4;
bytes reserved = 6; bytes reserved = 6;
DomainStrategy domain_strategy = 7;
bool is_client = 8; bool is_client = 8;
bool no_kernel_tun = 9; bool no_kernel_tun = 9;
repeated string DNS = 10; repeated string DNS = 10;
+14 -157
View File
@@ -15,8 +15,6 @@ import (
"net" "net"
"net/netip" "net/netip"
"os" "os"
"regexp"
"strconv"
"strings" "strings"
"syscall" "syscall"
"time" "time"
@@ -44,7 +42,6 @@ type netTun struct {
events chan tun.Event events chan tun.Event
notifyHandle *channel.NotificationHandle notifyHandle *channel.NotificationHandle
incomingPacket chan *buffer.View incomingPacket chan *buffer.View
closed chan struct{}
mtu int mtu int
dnsServers []netip.Addr dnsServers []netip.Addr
hasV4, hasV6 bool hasV4, hasV6 bool
@@ -61,7 +58,6 @@ func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int, handleLocal
stack: stack.New(opts), stack: stack.New(opts),
events: make(chan tun.Event, 10), events: make(chan tun.Event, 10),
incomingPacket: make(chan *buffer.View), incomingPacket: make(chan *buffer.View),
closed: make(chan struct{}),
dnsServers: dnsServers, dnsServers: dnsServers,
mtu: mtu, mtu: mtu,
} }
@@ -128,10 +124,8 @@ func (tun *netTun) Events() <-chan tun.Event {
} }
func (tun *netTun) Read(buf [][]byte, sizes []int, offset int) (int, error) { func (tun *netTun) Read(buf [][]byte, sizes []int, offset int) (int, error) {
var view *buffer.View view, ok := <-tun.incomingPacket
select { if !ok {
case view = <-tun.incomingPacket:
case <-tun.closed:
return 0, os.ErrClosed return 0, os.ErrClosed
} }
@@ -172,10 +166,7 @@ func (tun *netTun) WriteNotify() {
view := pkt.ToView() view := pkt.ToView()
pkt.DecRef() pkt.DecRef()
select { tun.incomingPacket <- view
case tun.incomingPacket <- view:
case <-tun.closed:
}
} }
func (tun *netTun) Close() error { func (tun *netTun) Close() error {
@@ -188,9 +179,8 @@ func (tun *netTun) Close() error {
close(tun.events) close(tun.events)
} }
// we don't close incomingPacket, because WriteNotify may be mid-send on it (DNS lookup) and would panic. if tun.incomingPacket != nil {
if tun.closed != nil { close(tun.incomingPacket)
close(tun.closed)
} }
return nil return nil
@@ -229,7 +219,6 @@ type Net struct {
DialUDPAddrPort func(laddr, raddr netip.AddrPort) (net.Conn, error) DialUDPAddrPort func(laddr, raddr netip.AddrPort) (net.Conn, error)
dnsServers []netip.Addr dnsServers []netip.Addr
hasV4, hasV6 bool hasV4, hasV6 bool
cache cache
} }
func convertToFullAddr(endpoint netip.AddrPort) (tcpip.FullAddress, tcpip.NetworkProtocolNumber) { func convertToFullAddr(endpoint netip.AddrPort) (tcpip.FullAddress, tcpip.NetworkProtocolNumber) {
@@ -257,12 +246,9 @@ var (
errServerTemporarilyMisbehaving = errors.New("server misbehaving") errServerTemporarilyMisbehaving = errors.New("server misbehaving")
errCanceled = errors.New("operation was canceled") errCanceled = errors.New("operation was canceled")
errTimeout = errors.New("i/o timeout") errTimeout = errors.New("i/o timeout")
errNumericPort = errors.New("port must be numeric")
errNoSuitableAddress = errors.New("no suitable address found")
errMissingAddress = errors.New("missing address")
) )
func (net *Net) LookupHost(host string) (addrs []string, err error) { func (net *Net) LookupHost(host string) (addrs []net.IP, ttl uint32, err error) {
return net.LookupContextHost(context.Background(), host) return net.LookupContextHost(context.Background(), host)
} }
@@ -581,12 +567,9 @@ func (tnet *Net) tryOneName(ctx context.Context, name string, qtype dnsmessage.T
return dnsmessage.Parser{}, "", lastErr return dnsmessage.Parser{}, "", lastErr
} }
func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, error) { func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]net.IP, uint32, error) {
if saddr := tnet.cache.LookupHost(host); saddr != nil {
return saddr, nil
}
if host == "" || (!tnet.hasV6 && !tnet.hasV4) { if host == "" || (!tnet.hasV6 && !tnet.hasV4) {
return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} return nil, 0, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
} }
zlen := len(host) zlen := len(host)
if strings.IndexByte(host, ':') != -1 { if strings.IndexByte(host, ':') != -1 {
@@ -595,11 +578,11 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string,
} }
} }
if ip, err := netip.ParseAddr(host[:zlen]); err == nil { if ip, err := netip.ParseAddr(host[:zlen]); err == nil {
return []string{ip.String()}, nil return []net.IP{ip.AsSlice()}, 0, nil
} }
if !isDomainName(host) { if !isDomainName(host) {
return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} return nil, 0, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
} }
type result struct { type result struct {
p dnsmessage.Parser p dnsmessage.Parser
@@ -700,137 +683,11 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string,
} }
if len(addrs) == 0 && lastErr != nil { if len(addrs) == 0 && lastErr != nil {
return nil, lastErr return nil, 0, lastErr
} }
saddrs := make([]string, 0, len(addrs)) ips := make([]net.IP, 0, len(addrs))
for _, ip := range addrs { for _, ip := range addrs {
saddrs = append(saddrs, ip.String()) ips = append(ips, ip.AsSlice())
} }
tnet.cache.Cache(host, saddrs, ttl) return ips, ttl, nil
return saddrs, nil
}
func partialDeadline(now, deadline time.Time, addrsRemaining int) (time.Time, error) {
if deadline.IsZero() {
return deadline, nil
}
timeRemaining := deadline.Sub(now)
if timeRemaining <= 0 {
return time.Time{}, errTimeout
}
timeout := timeRemaining / time.Duration(addrsRemaining)
const saneMinimum = 2 * time.Second
if timeout < saneMinimum {
if timeRemaining < saneMinimum {
timeout = timeRemaining
} else {
timeout = saneMinimum
}
}
return now.Add(timeout), nil
}
var protoSplitter = regexp.MustCompile(`^(tcp|udp|ping)(4|6)?$`)
func (tnet *Net) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
if ctx == nil {
panic("nil context")
}
var acceptV4, acceptV6 bool
matches := protoSplitter.FindStringSubmatch(network)
if matches == nil {
return nil, &net.OpError{Op: "dial", Err: net.UnknownNetworkError(network)}
} else if len(matches[2]) == 0 {
acceptV4 = true
acceptV6 = true
} else {
acceptV4 = matches[2][0] == '4'
acceptV6 = !acceptV4
}
var host string
var port int
if matches[1] == "ping" {
host = address
} else {
var sport string
var err error
host, sport, err = net.SplitHostPort(address)
if err != nil {
return nil, &net.OpError{Op: "dial", Err: err}
}
port, err = strconv.Atoi(sport)
if err != nil || port < 0 || port > 65535 {
return nil, &net.OpError{Op: "dial", Err: errNumericPort}
}
}
allAddr, err := tnet.LookupContextHost(ctx, host)
if err != nil {
return nil, &net.OpError{Op: "dial", Err: err}
}
var addrs []netip.AddrPort
for _, addr := range allAddr {
ip, err := netip.ParseAddr(addr)
if err == nil && ((ip.Is4() && acceptV4) || (ip.Is6() && acceptV6)) {
addrs = append(addrs, netip.AddrPortFrom(ip, uint16(port)))
}
}
if len(addrs) == 0 && len(allAddr) != 0 {
return nil, &net.OpError{Op: "dial", Err: errNoSuitableAddress}
}
var firstErr error
for i, addr := range addrs {
select {
case <-ctx.Done():
err := ctx.Err()
if err == context.Canceled {
err = errCanceled
} else if err == context.DeadlineExceeded {
err = errTimeout
}
return nil, &net.OpError{Op: "dial", Err: err}
default:
}
dialCtx := ctx
if deadline, hasDeadline := ctx.Deadline(); hasDeadline {
partialDeadline, err := partialDeadline(time.Now(), deadline, len(addrs)-i)
if err != nil {
if firstErr == nil {
firstErr = &net.OpError{Op: "dial", Err: err}
}
break
}
if partialDeadline.Before(deadline) {
var cancel context.CancelFunc
dialCtx, cancel = context.WithDeadline(ctx, partialDeadline)
defer cancel()
}
}
var c net.Conn
switch matches[1] {
case "tcp":
c, err = tnet.DialContextTCPAddrPort(dialCtx, addr)
case "udp":
c, err = tnet.DialUDPAddrPort(netip.AddrPort{}, addr)
case "ping":
err = errors.New("not support")
// c, err = tnet.DialPingAddr(netip.Addr{}, addr.Addr())
}
if err == nil {
return c, nil
}
if firstErr == nil {
firstErr = err
}
}
if firstErr == nil {
firstErr = &net.OpError{Op: "dial", Err: errMissingAddress}
}
return nil, firstErr
}
func (tnet *Net) Dial(network, address string) (net.Conn, error) {
return tnet.DialContext(context.Background(), network, address)
} }
+9 -7
View File
@@ -258,16 +258,18 @@ func (s *Server) Start() error {
return errors.New("address is domain") return errors.New("address is domain")
} }
listenFunc := func() (net.PacketConn, error) { listenFunc := func() (net.PacketConn, error) {
var pktConn net.PacketConn pktConn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)}, s.streamSettings.SocketSettings)
var err error
if s.streamSettings.FinalMask != nil {
pktConn, err = s.streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)})
} else {
pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)}, s.streamSettings.SocketSettings)
}
if err != nil { if err != nil {
return nil, err return nil, err
} }
if s.streamSettings.UdpmaskManager != nil {
newConn, err := s.streamSettings.UdpmaskManager.WrapPacketConnServer(pktConn)
if err != nil {
pktConn.Close()
return nil, errors.New("mask err").Base(err)
}
pktConn = newConn
}
if s.uplinkCounter != nil || s.downlinkCounter != nil { if s.uplinkCounter != nil || s.downlinkCounter != nil {
pktConn = &PacketCounterConnection{ pktConn = &PacketCounterConnection{
PacketConn: pktConn, PacketConn: pktConn,
-2
View File
@@ -65,7 +65,6 @@ func TestWireguard(t *testing.T) {
ProxySettings: serial.ToTypedMessage(&freedom.Config{ ProxySettings: serial.ToTypedMessage(&freedom.Config{
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}}, FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
}), }),
SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{}),
}, },
}, },
} }
@@ -105,7 +104,6 @@ func TestWireguard(t *testing.T) {
AllowedIps: []string{"0.0.0.0/0", "::0/0"}, AllowedIps: []string{"0.0.0.0/0", "::0/0"},
}}, }},
}), }),
SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{}),
}, },
}, },
} }
+1 -1
View File
@@ -27,7 +27,7 @@ var strategy = [11][3]byte{
func RegisterProtocolConfigCreator(name string, creator ConfigCreator) error { func RegisterProtocolConfigCreator(name string, creator ConfigCreator) error {
if _, found := globalTransportConfigCreatorCache[name]; found { if _, found := globalTransportConfigCreatorCache[name]; found {
return errors.New("protocol ", name, " is already registered").AtError() return errors.New("protocol ", name, " is already registered")
} }
globalTransportConfigCreatorCache[name] = creator globalTransportConfigCreatorCache[name] = creator
return nil return nil
+6 -6
View File
@@ -38,7 +38,7 @@ var transportDialerCache = make(map[string]dialFunc)
// RegisterTransportDialer registers a Dialer with given name. // RegisterTransportDialer registers a Dialer with given name.
func RegisterTransportDialer(protocol string, dialer dialFunc) error { func RegisterTransportDialer(protocol string, dialer dialFunc) error {
if _, found := transportDialerCache[protocol]; found { if _, found := transportDialerCache[protocol]; found {
return errors.New(protocol, " dialer already registered").AtError() return errors.New(protocol, " dialer already registered")
} }
transportDialerCache[protocol] = dialer transportDialerCache[protocol] = dialer
return nil return nil
@@ -58,7 +58,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *MemoryStrea
protocol := streamSettings.ProtocolName protocol := streamSettings.ProtocolName
dialer := transportDialerCache[protocol] dialer := transportDialerCache[protocol]
if dialer == nil { if dialer == nil {
return nil, errors.New(protocol, " dialer not registered").AtError() return nil, errors.New(protocol, " dialer not registered")
} }
return dialer(ctx, dest, streamSettings) return dialer(ctx, dest, streamSettings)
} }
@@ -66,7 +66,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *MemoryStrea
if dest.Network == net.Network_UDP { if dest.Network == net.Network_UDP {
udpDialer := transportDialerCache["udp"] udpDialer := transportDialerCache["udp"]
if udpDialer == nil { if udpDialer == nil {
return nil, errors.New("UDP dialer not registered").AtError() return nil, errors.New("UDP dialer not registered")
} }
return udpDialer(ctx, dest, streamSettings) return udpDialer(ctx, dest, streamSettings)
} }
@@ -86,7 +86,7 @@ var (
func LookupForIP(domain string, strategy DomainStrategy, localAddr net.Address) ([]net.IP, error) { func LookupForIP(domain string, strategy DomainStrategy, localAddr net.Address) ([]net.IP, error) {
if dnsClient == nil { if dnsClient == nil {
return nil, errors.New("DNS client not initialized").AtError() return nil, errors.New("DNS client not initialized")
} }
ips, _, err := dnsClient.LookupIP(domain, dns.IPOption{ ips, _, err := dnsClient.LookupIP(domain, dns.IPOption{
@@ -269,11 +269,11 @@ func DialSystem(ctx context.Context, dest net.Destination, sockopt *SocketConfig
if len(sockopt.DialerProxy) > 0 { if len(sockopt.DialerProxy) > 0 {
if obm == nil { if obm == nil {
return nil, errors.New("there is no outbound manager for dialerProxy").AtError() return nil, errors.New("there is no outbound manager for dialerProxy")
} }
h := obm.GetHandler(sockopt.DialerProxy) h := obm.GetHandler(sockopt.DialerProxy)
if h == nil { if h == nil {
return nil, errors.New("there is no outbound handler for dialerProxy").AtError() return nil, errors.New("there is no outbound handler for dialerProxy")
} }
return redirect(ctx, dest, sockopt.DialerProxy, h), nil return redirect(ctx, dest, sockopt.DialerProxy, h), nil
} }
+95 -238
View File
@@ -2,291 +2,103 @@ package finalmask
import ( import (
"context" "context"
"fmt" "net"
"slices" "slices"
"github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
) )
type Dialer struct { type Udpmask interface {
DialTCP func(net.Destination) (net.Conn, error) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error)
DialUDP func(net.Destination) (net.Conn, error) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error)
} }
type ListenConfig struct { type UdpmaskManager struct {
Listen func(net.Addr) (net.Listener, error) udpmasks []Udpmask
ListenPacket func(net.Addr) (net.PacketConn, error)
} }
type TCPMask interface { func NewUdpmaskManager(udpmasks []Udpmask) *UdpmaskManager {
WrapConnClient(net.Conn, *net.Destination, *Dialer) (net.Conn, error) slices.Reverse(udpmasks)
WrapConnServer(net.Conn) (net.Conn, error) return &UdpmaskManager{udpmasks: udpmasks}
// Listen(net.Listener) (net.Listener, error)
} }
type UDPMask interface { func (m *UdpmaskManager) WrapPacketConnClient(raw net.PacketConn) (net.PacketConn, error) {
WrapPacketConnClient(net.PacketConn, *net.Destination, *Dialer) (net.PacketConn, error)
WrapPacketConnServer(net.PacketConn, net.Addr, *ListenConfig) (net.PacketConn, error)
}
type FinalMask struct {
tcpMasks []TCPMask
udpMasks []UDPMask
dialTCP func(context.Context, net.Destination) (net.Conn, error)
listen func(context.Context, net.Addr) (net.Listener, error)
dialUDP func(context.Context, net.Destination) (net.PacketConn, net.Addr, error)
listenPacket func(context.Context, net.Addr) (net.PacketConn, error)
}
func NewFinalMask(tcpMasks []TCPMask, udpMasks []UDPMask, dialTCP func(context.Context, net.Destination) (net.Conn, error), listen func(context.Context, net.Addr) (net.Listener, error), dialUDP func(context.Context, net.Destination) (net.PacketConn, net.Addr, error), listenPacket func(context.Context, net.Addr) (net.PacketConn, error)) *FinalMask {
slices.Reverse(tcpMasks)
slices.Reverse(udpMasks)
return &FinalMask{
tcpMasks: tcpMasks,
udpMasks: udpMasks,
dialTCP: dialTCP,
dialUDP: dialUDP,
listen: listen,
listenPacket: listenPacket,
}
}
func (fm *FinalMask) DialTCP(ctx context.Context, dest net.Destination) (net.Conn, error) {
if len(fm.tcpMasks) == 0 {
return fm.dialTCP(ctx, dest)
}
for i := range fm.tcpMasks {
if i > 0 {
if _, ok := fm.tcpMasks[i].(interface{ HandleDial() }); ok {
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.tcpMasks[i])
}
}
}
var conn net.Conn
var err error
if _, ok := fm.tcpMasks[0].(interface{ HandleDial() }); !ok {
conn, err = fm.dialTCP(ctx, dest)
if err != nil {
return nil, err
}
}
dialer := &Dialer{
DialTCP: func(dest net.Destination) (net.Conn, error) {
return fm.dialTCP(ctx, dest)
},
DialUDP: func(dest net.Destination) (net.Conn, error) {
conn, addr, err := fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
},
}
for i := range fm.tcpMasks {
var newConn net.Conn
newConn, err = fm.tcpMasks[i].WrapConnClient(conn, &dest, dialer)
if err != nil {
_ = conn.Close()
return nil, err
}
conn = newConn
}
return conn, nil
}
func (fm *FinalMask) Listen(ctx context.Context, addr net.Addr) (net.Listener, error) {
if len(fm.tcpMasks) == 0 {
return fm.listen(ctx, addr)
}
off := 0
listener, err := fm.listen(ctx, addr)
if err != nil {
return nil, err
}
for i := range fm.tcpMasks {
if _, ok := fm.tcpMasks[i].(interface {
Listen(net.Listener) (net.Listener, error)
}); ok {
if i-off == 0 {
l, err := fm.tcpMasks[i].(interface {
Listen(net.Listener) (net.Listener, error)
}).Listen(listener)
if err != nil {
listener.Close()
return nil, err
}
listener = l
} else {
l, err := fm.tcpMasks[i].(interface {
Listen(net.Listener) (net.Listener, error)
}).Listen(&TCPListener{Listener: listener, tcpMasks: fm.tcpMasks[off:i]})
if err != nil {
listener.Close()
return nil, err
}
listener = l
}
off = i + 1
}
}
if off < len(fm.tcpMasks) {
return &TCPListener{Listener: listener, tcpMasks: fm.tcpMasks[off:]}, nil
}
return listener, nil
}
func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Conn, error) {
if len(fm.udpMasks) == 0 {
conn, addr, err := fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
}
for i := range fm.udpMasks {
if i > 0 {
if _, ok := fm.udpMasks[i].(interface{ HandleDial() }); ok {
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.udpMasks[i])
}
}
}
var conn net.PacketConn
var addr net.Addr
var err error
if _, ok := fm.udpMasks[0].(interface{ HandleDial() }); !ok {
conn, addr, err = fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
}
dialer := &Dialer{
DialTCP: func(dest net.Destination) (net.Conn, error) {
return fm.dialTCP(ctx, dest)
},
DialUDP: func(dest net.Destination) (net.Conn, error) {
conn, addr, err := fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
},
}
var sizes []int var sizes []int
var conns []net.PacketConn var conns []net.PacketConn
for i := range fm.udpMasks { for i, mask := range m.udpmasks {
var newConn net.PacketConn if _, ok := mask.(headerConn); ok {
if _, ok := fm.udpMasks[i].(interface{ HeaderConn() }); ok { conn, err := mask.WrapPacketConnClient(nil, i, len(m.udpmasks)-1)
newConn, err = fm.udpMasks[i].WrapPacketConnClient(nil, nil, nil)
if err != nil { if err != nil {
_ = conn.Close()
return nil, err return nil, err
} }
sizes = append(sizes, newConn.(interface{ Size() int }).Size()) sizes = append(sizes, conn.(headerSize).Size())
conns = append(conns, newConn) conns = append(conns, conn)
} else { } else {
if len(conns) > 0 { if len(conns) > 0 {
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns} raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
sizes = nil sizes = nil
conns = nil conns = nil
} }
newConn, err = fm.udpMasks[i].WrapPacketConnClient(conn, &dest, dialer) var err error
raw, err = mask.WrapPacketConnClient(raw, i, len(m.udpmasks)-1)
if err != nil { if err != nil {
_ = conn.Close()
return nil, err return nil, err
} }
conn = newConn
} }
} }
if len(conns) > 0 { if len(conns) > 0 {
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns} raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
sizes = nil sizes = nil
conns = nil conns = nil
} }
if addr == nil { return raw, nil
addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
} }
func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) { func (m *UdpmaskManager) WrapPacketConnServer(raw net.PacketConn) (net.PacketConn, error) {
if len(fm.udpMasks) == 0 {
return fm.listenPacket(ctx, addr)
}
for i := range fm.udpMasks {
if i > 0 {
if _, ok := fm.udpMasks[i].(interface{ HandleListen() }); ok {
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.udpMasks[i])
}
}
}
var conn net.PacketConn
var err error
if _, ok := fm.udpMasks[0].(interface{ HandleListen() }); !ok {
conn, err = fm.listenPacket(ctx, addr)
if err != nil {
return nil, err
}
}
lc := &ListenConfig{
Listen: func(addr net.Addr) (net.Listener, error) { return fm.listen(ctx, addr) },
ListenPacket: func(addr net.Addr) (net.PacketConn, error) { return fm.listenPacket(ctx, addr) },
}
var sizes []int var sizes []int
var conns []net.PacketConn var conns []net.PacketConn
for i := range fm.udpMasks { for i, mask := range m.udpmasks {
var newConn net.PacketConn if _, ok := mask.(headerConn); ok {
if _, ok := fm.udpMasks[i].(interface{ HeaderConn() }); ok { conn, err := mask.WrapPacketConnServer(nil, i, len(m.udpmasks)-1)
newConn, err = fm.udpMasks[i].WrapPacketConnServer(nil, nil, nil)
if err != nil { if err != nil {
_ = conn.Close()
return nil, err return nil, err
} }
sizes = append(sizes, newConn.(interface{ Size() int }).Size()) sizes = append(sizes, conn.(headerSize).Size())
conns = append(conns, newConn) conns = append(conns, conn)
} else { } else {
if len(conns) > 0 { if len(conns) > 0 {
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns} raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
sizes = nil sizes = nil
conns = nil conns = nil
} }
newConn, err = fm.udpMasks[i].WrapPacketConnServer(conn, addr, lc) var err error
raw, err = mask.WrapPacketConnServer(raw, i, len(m.udpmasks)-1)
if err != nil { if err != nil {
_ = conn.Close()
return nil, err return nil, err
} }
conn = newConn
} }
} }
if len(conns) > 0 { if len(conns) > 0 {
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns} raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
sizes = nil sizes = nil
conns = nil conns = nil
} }
return conn, nil return raw, nil
} }
const ( const (
UDPSize = 4096 UDPSize = 4096
) )
type PacketConnWrapper struct { type headerConn interface {
net.PacketConn HeaderConn()
udpAddr net.Addr
} }
func (c *PacketConnWrapper) RemoteAddr() net.Addr { type headerSize interface {
return c.udpAddr Size() int
}
func (c *PacketConnWrapper) Read(b []byte) (n int, err error) {
n, _, err = c.PacketConn.ReadFrom(b)
return
}
func (c *PacketConnWrapper) Write(b []byte) (n int, err error) {
return c.PacketConn.WriteTo(b, c.udpAddr)
} }
type headerManagerConn struct { type headerManagerConn struct {
@@ -379,27 +191,72 @@ func (c *headerManagerConn) WriteTo(p []byte, addr net.Addr) (n int, err error)
return len(p), nil return len(p), nil
} }
type TCPListener struct { type Tcpmask interface {
net.Listener WrapConnClient(net.Conn) (net.Conn, error)
tcpMasks []TCPMask WrapConnServer(net.Conn) (net.Conn, error)
} }
func (l *TCPListener) Accept() (net.Conn, error) { type TcpmaskManager struct {
tcpmasks []Tcpmask
}
func NewTcpmaskManager(tcpmasks []Tcpmask) *TcpmaskManager {
slices.Reverse(tcpmasks)
return &TcpmaskManager{tcpmasks: tcpmasks}
}
func (m *TcpmaskManager) WrapConnClient(raw net.Conn) (net.Conn, error) {
var err error
for _, mask := range m.tcpmasks {
raw, err = mask.WrapConnClient(raw)
if err != nil {
return nil, err
}
}
return raw, nil
}
func (m *TcpmaskManager) WrapConnServer(raw net.Conn) (net.Conn, error) {
var err error
for _, mask := range m.tcpmasks {
raw, err = mask.WrapConnServer(raw)
if err != nil {
return nil, err
}
}
return raw, nil
}
func (m *TcpmaskManager) WrapListener(l net.Listener) (net.Listener, error) {
return NewTcpListener(m, l)
}
type tcpListener struct {
m *TcpmaskManager
net.Listener
}
func NewTcpListener(m *TcpmaskManager, l net.Listener) (net.Listener, error) {
return &tcpListener{
m: m,
Listener: l,
}, nil
}
func (l *tcpListener) Accept() (net.Conn, error) {
conn, err := l.Listener.Accept() conn, err := l.Listener.Accept()
if err != nil { if err != nil {
return conn, err return conn, err
} }
for i := range l.tcpMasks { newConn, err := l.m.WrapConnServer(conn)
var newConn net.Conn if err != nil {
newConn, err = l.tcpMasks[i].WrapConnServer(conn) errors.LogDebugInner(context.Background(), err, "mask err")
if err != nil { _ = conn.Close()
_ = conn.Close() return nil, err
return nil, err
}
conn = newConn
} }
return conn, nil
return newConn, nil
} }
type TcpMaskConn interface { type TcpMaskConn interface {
@@ -1,14 +1,11 @@
package fragment package fragment
import ( import "net"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) { func (c *Config) WrapConnClient(raw net.Conn) (net.Conn, error) {
return NewConnClient(c, conn, false) return NewConnClient(c, raw, false)
} }
func (c *Config) WrapConnServer(conn net.Conn) (net.Conn, error) { func (c *Config) WrapConnServer(raw net.Conn) (net.Conn, error) {
return NewConnServer(c, conn, true) return NewConnServer(c, raw, true)
} }
@@ -1,30 +1,29 @@
package custom package custom
import ( import (
"github.com/xtls/xray-core/common/net" "net"
"github.com/xtls/xray-core/transport/internet/finalmask"
) )
func (c *TCPConfig) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) { func (c *TCPConfig) WrapConnClient(raw net.Conn) (net.Conn, error) {
return NewConnClientTCP(c, conn) return NewConnClientTCP(c, raw)
} }
func (c *TCPConfig) WrapConnServer(conn net.Conn) (net.Conn, error) { func (c *TCPConfig) WrapConnServer(raw net.Conn) (net.Conn, error) {
return NewConnServerTCP(c, conn) return NewConnServerTCP(c, raw)
} }
func (c *UDPConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *UDPConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClientUDP(c, conn) return NewConnClientUDP(c, raw)
} }
func (c *UDPConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *UDPConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServerUDP(c, conn) return NewConnServerUDP(c, raw)
} }
func (c *UDPStandaloneConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *UDPStandaloneConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClientUDPStandalone(c, conn) return NewConnClientUDPStandalone(c, raw)
} }
func (c *UDPStandaloneConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *UDPStandaloneConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServerUDPStandalone(c, conn) return NewConnServerUDPStandalone(c, raw)
} }
@@ -9,6 +9,8 @@ import (
"strings" "strings"
"testing" "testing"
"time" "time"
"github.com/xtls/xray-core/transport/internet/finalmask"
) )
func TestMetadataEvaluatorRejectsUnknownName(t *testing.T) { func TestMetadataEvaluatorRejectsUnknownName(t *testing.T) {
@@ -154,7 +156,7 @@ func TestMetadataUDPStandaloneWriteUsesRemotePort(t *testing.T) {
} }
defer serverRaw.Close() defer serverRaw.Close()
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil) client, err := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg}).WrapPacketConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -299,7 +301,7 @@ func TestMetadataTCPHandshakeUsesEndpointPorts(t *testing.T) {
} }
defer serverRaw.Close() defer serverRaw.Close()
client, err := clientCfg.WrapConnClient(clientRaw, nil, nil) client, err := clientCfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -5,6 +5,8 @@ import (
"net" "net"
"testing" "testing"
"time" "time"
"github.com/xtls/xray-core/transport/internet/finalmask"
) )
func mustSendRecvUDP(t *testing.T, from net.PacketConn, to net.PacketConn, msg []byte) { func mustSendRecvUDP(t *testing.T, from net.PacketConn, to net.PacketConn, msg []byte) {
@@ -46,6 +48,7 @@ func TestStateUDPResponseReusesPriorCapturedValues(t *testing.T) {
}, },
}, },
} }
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0") clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil { if err != nil {
@@ -59,11 +62,11 @@ func TestStateUDPResponseReusesPriorCapturedValues(t *testing.T) {
} }
defer serverRaw.Close() defer serverRaw.Close()
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil) client, err := maskManager.WrapPacketConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil) server, err := maskManager.WrapPacketConnServer(serverRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -37,7 +37,7 @@ func TestDSLTCPHandshakeReusesCapturedValue(t *testing.T) {
defer clientRaw.Close() defer clientRaw.Close()
defer serverRaw.Close() defer serverRaw.Close()
client, err := cfg.WrapConnClient(clientRaw, nil, nil) client, err := cfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -117,7 +117,7 @@ func TestDSLTCPClientRejectsMismatchedResponseSequence(t *testing.T) {
defer clientRaw.Close() defer clientRaw.Close()
defer serverRaw.Close() defer serverRaw.Close()
client, err := clientCfg.WrapConnClient(clientRaw, nil, nil) client, err := clientCfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -1,16 +1,15 @@
package aes128gcm package aes128gcm
import ( import (
"github.com/xtls/xray-core/common/net" "net"
"github.com/xtls/xray-core/transport/internet/finalmask"
) )
func (c *Config) HeaderConn() {} func (c *Config) HeaderConn() {}
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClient(c, conn) return NewConnClient(c, raw)
} }
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServer(c, conn) return NewConnServer(c, raw)
} }
@@ -1,16 +1,15 @@
package header package header
import ( import (
"github.com/xtls/xray-core/common/net" "net"
"github.com/xtls/xray-core/transport/internet/finalmask"
) )
func (c *Config) HeaderConn() {} func (c *Config) HeaderConn() {}
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClient(c, conn) return NewConnClient(c, raw)
} }
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServer(c, conn) return NewConnServer(c, raw)
} }
@@ -1,16 +1,15 @@
package original package original
import ( import (
"github.com/xtls/xray-core/common/net" "net"
"github.com/xtls/xray-core/transport/internet/finalmask"
) )
func (c *Config) HeaderConn() {} func (c *Config) HeaderConn() {}
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClient(c, conn) return NewConnClient(c, raw)
} }
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServer(c, conn) return NewConnServer(c, raw)
} }
+5 -8
View File
@@ -1,14 +1,11 @@
package noise package noise
import ( import "net"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClient(c, conn) return NewConnClient(c, raw)
} }
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServer(c, conn) return NewConnServer(c, raw)
} }
+15 -6
View File
@@ -1,14 +1,23 @@
package realm package realm
import ( import (
"github.com/xtls/xray-core/common/net" "net"
"github.com/xtls/xray-core/transport/internet/finalmask"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet"
) )
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClient(c, conn) _, ok1 := raw.(*internet.FakePacketConn)
if level != 0 || ok1 {
return nil, errors.New("realm requires being at the outermost level")
}
return NewConnClient(c, raw)
} }
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServer(c, conn) if level != 0 {
return nil, errors.New("realm requires being at the outermost level")
}
return NewConnServer(c, raw)
} }
@@ -1,24 +1,23 @@
package salamander package salamander
import ( import (
"github.com/xtls/xray-core/common/net" "net"
"github.com/xtls/xray-core/transport/internet/finalmask"
) )
func (c *Config) HeaderConn() {} func (c *Config) HeaderConn() {}
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewSalamanderConnClient(c, conn) return NewSalamanderConnClient(c, raw)
} }
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewSalamanderConnServer(c, conn) return NewSalamanderConnServer(c, raw)
} }
func (c *GeckoConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *GeckoConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewGeckoConnClient(c, conn) return NewGeckoConnClient(c, raw)
} }
func (c *GeckoConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *GeckoConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewGeckoConnServer(c, conn) return NewGeckoConnServer(c, raw)
} }
+17 -10
View File
@@ -1,18 +1,19 @@
package sudoku package sudoku
import ( import (
"github.com/xtls/xray-core/common/net" "net"
"github.com/xtls/xray-core/transport/internet/finalmask"
"github.com/xtls/xray-core/common/errors"
) )
// Sudoku in finalmask mode is a pure appearance transform with no standalone handshake. // Sudoku in finalmask mode is a pure appearance transform with no standalone handshake.
// TCP always keeps classic sudoku on uplink and uses packed downlink optimization on server writes. // TCP always keeps classic sudoku on uplink and uses packed downlink optimization on server writes.
func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) { func (c *Config) WrapConnClient(raw net.Conn) (net.Conn, error) {
return newPackedDirectionalConn(conn, c, true) return newPackedDirectionalConn(raw, c, true)
} }
func (c *Config) WrapConnServer(conn net.Conn) (net.Conn, error) { func (c *Config) WrapConnServer(raw net.Conn) (net.Conn, error) {
return newPackedDirectionalConn(conn, c, false) return newPackedDirectionalConn(raw, c, false)
} }
func newPackedDirectionalConn(raw net.Conn, config *Config, readPacked bool) (net.Conn, error) { func newPackedDirectionalConn(raw net.Conn, config *Config, readPacked bool) (net.Conn, error) {
@@ -35,10 +36,16 @@ func newPackedDirectionalConn(raw net.Conn, config *Config, readPacked bool) (ne
return newWrappedConn(raw, reader, writer), nil return newWrappedConn(raw, reader, writer), nil
} }
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewUDPConn(conn, c) if level != levelCount {
return nil, errors.New("sudoku udp mask must be the innermost mask in chain")
}
return NewUDPConn(raw, c)
} }
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewUDPConn(conn, c) if level != levelCount {
return nil, errors.New("sudoku udp mask must be the innermost mask in chain")
}
return NewUDPConn(raw, c)
} }
+55 -57
View File
@@ -2,14 +2,12 @@ package finalmask_test
import ( import (
"bytes" "bytes"
"context"
"io" "io"
gonet "net" "net"
"strings" "strings"
"testing" "testing"
"time" "time"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet/finalmask"
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom" "github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
) )
@@ -22,14 +20,11 @@ func mustSendRecvTcp(
) { ) {
t.Helper() t.Helper()
waitCh := make(chan error)
go func() { go func() {
_, err := from.Write(msg) _, err := from.Write(msg)
if err != nil { if err != nil {
t.Fatal(err) t.Error(err)
} }
close(waitCh)
}() }()
buf := make([]byte, 1024) buf := make([]byte, 1024)
@@ -45,23 +40,18 @@ func mustSendRecvTcp(
if !bytes.Equal(buf[:n], msg) { if !bytes.Equal(buf[:n], msg) {
t.Fatalf("unexpected data %q", buf[:n]) t.Fatalf("unexpected data %q", buf[:n])
} }
<-waitCh
} }
type layerMaskTcp struct { type layerMaskTcp struct {
name string name string
mask finalmask.TCPMask mask finalmask.Tcpmask
} }
type failingWrapMask struct{} type failingWrapMask struct{}
func (failingWrapMask) TCP() {} func (failingWrapMask) TCP() {}
func (f failingWrapMask) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) { func (f failingWrapMask) WrapConnClient(raw net.Conn) (net.Conn, error) { return raw, nil }
return conn, nil func (f failingWrapMask) WrapConnServer(raw net.Conn) (net.Conn, error) {
}
func (f failingWrapMask) WrapConnServer(conn net.Conn) (net.Conn, error) {
return nil, io.ErrClosedPipe return nil, io.ErrClosedPipe
} }
@@ -102,31 +92,32 @@ func TestConnReadWrite(t *testing.T) {
t.Run(c.name, func(t *testing.T) { t.Run(c.name, func(t *testing.T) {
mask := c.mask mask := c.mask
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) { maskManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{mask})
return net.Dial("tcp", dest.NetAddr())
}
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
return net.Listen("tcp", addr.String())
}
finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{mask}, nil, dialTCP, listen, nil, nil)
listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()}) ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
t.Cleanup(func() { listener.Close() })
client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port))) client, err := net.Dial("tcp", ln.Addr().String())
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
t.Cleanup(func() { client.Close() })
server, err := listener.Accept() client, err = maskManager.WrapConnClient(client)
if err != nil {
t.Fatal(err)
}
server, err := ln.Accept()
if err != nil {
t.Fatal(err)
}
server, err = maskManager.WrapConnServer(server)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
t.Cleanup(func() { server.Close() })
_ = client.SetDeadline(time.Now().Add(time.Second)) _ = client.SetDeadline(time.Now().Add(time.Second))
_ = server.SetDeadline(time.Now().Add(time.Second)) _ = server.SetDeadline(time.Now().Add(time.Second))
@@ -159,32 +150,34 @@ func TestTCPcustomStaticHandshakeRoundTrip(t *testing.T) {
}, },
}, },
} }
maskManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{cfg})
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) { ln, err := net.Listen("tcp", "127.0.0.1:0")
return net.Dial("tcp", dest.NetAddr())
}
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
return net.Listen("tcp", addr.String())
}
finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{cfg}, nil, dialTCP, listen, nil, nil)
listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()})
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
defer listener.Close() defer ln.Close()
client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port))) clientRaw, err := net.Dial("tcp", ln.Addr().String())
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
defer client.Close() defer clientRaw.Close()
server, err := listener.Accept() serverRaw, err := ln.Accept()
if err != nil {
t.Fatal(err)
}
defer serverRaw.Close()
client, err := maskManager.WrapConnClient(clientRaw)
if err != nil {
t.Fatal(err)
}
server, err := maskManager.WrapConnServer(serverRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
defer server.Close()
_ = client.SetDeadline(time.Now().Add(time.Second)) _ = client.SetDeadline(time.Now().Add(time.Second))
_ = server.SetDeadline(time.Now().Add(time.Second)) _ = server.SetDeadline(time.Now().Add(time.Second))
@@ -227,11 +220,11 @@ func TestTCPcustomClientRejectsMismatchedServerSequence(t *testing.T) {
}, },
} }
clientRaw, serverRaw := gonet.Pipe() clientRaw, serverRaw := net.Pipe()
defer clientRaw.Close() defer clientRaw.Close()
defer serverRaw.Close() defer serverRaw.Close()
client, err := clientCfg.WrapConnClient(clientRaw, nil, nil) client, err := clientCfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -264,37 +257,42 @@ func TestTCPcustomClientRejectsMismatchedServerSequence(t *testing.T) {
} }
func TestTCPWrapListenerRejectsImmediateWrapErrors(t *testing.T) { func TestTCPWrapListenerRejectsImmediateWrapErrors(t *testing.T) {
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) { clientManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{failingWrapMask{}})
return net.Dial("tcp", dest.NetAddr()) serverManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{failingWrapMask{}})
}
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
return net.Listen("tcp", addr.String())
}
finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{failingWrapMask{}}, nil, dialTCP, listen, nil, nil)
listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()}) rawLn, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer rawLn.Close()
ln, err := serverManager.WrapListener(rawLn)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
defer listener.Close()
accepted := make(chan struct { accepted := make(chan struct {
conn net.Conn conn net.Conn
err error err error
}, 1) }, 1)
go func() { go func() {
conn, err := listener.Accept() conn, err := ln.Accept()
accepted <- struct { accepted <- struct {
conn net.Conn conn net.Conn
err error err error
}{conn: conn, err: err} }{conn: conn, err: err}
}() }()
client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port))) clientRaw, err := net.Dial("tcp", rawLn.Addr().String())
if err != nil {
t.Fatal(err)
}
defer clientRaw.Close()
client, err := clientManager.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
defer client.Close()
_ = client.SetDeadline(time.Now().Add(time.Second)) _ = client.SetDeadline(time.Now().Add(time.Second))
+59 -52
View File
@@ -2,15 +2,13 @@ package finalmask_test
import ( import (
"bytes" "bytes"
"context"
"encoding/binary" "encoding/binary"
"io" "io"
gonet "net" "net"
"sync/atomic" "sync/atomic"
"testing" "testing"
"time" "time"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/proxy" "github.com/xtls/xray-core/proxy"
"github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet/finalmask"
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom" "github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
@@ -53,7 +51,7 @@ func mustSendRecv(
type layerMask struct { type layerMask struct {
name string name string
mask finalmask.UDPMask mask finalmask.Udpmask
layers int layers int
} }
@@ -215,23 +213,25 @@ func newStandaloneStunLikeUDPServerConfig() *custom.UDPStandaloneConfig {
func newUDPClientServerPair(t *testing.T, cfg *custom.UDPStandaloneConfig) (net.PacketConn, net.PacketConn, net.PacketConn, net.PacketConn) { func newUDPClientServerPair(t *testing.T, cfg *custom.UDPStandaloneConfig) (net.PacketConn, net.PacketConn, net.PacketConn, net.PacketConn) {
t.Helper() t.Helper()
clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
t.Cleanup(func() { _ = clientRaw.Close() }) t.Cleanup(func() { _ = clientRaw.Close() })
serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
t.Cleanup(func() { _ = serverRaw.Close() }) t.Cleanup(func() { _ = serverRaw.Close() })
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil) maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
client, err := maskManager.WrapPacketConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil) server, err := maskManager.WrapPacketConnServer(serverRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -348,39 +348,31 @@ func TestPacketConnReadWrite(t *testing.T) {
if layers <= 0 { if layers <= 0 {
layers = 1 layers = 1
} }
masks := make([]finalmask.UDPMask, 0, layers) masks := make([]finalmask.Udpmask, 0, layers)
for i := 0; i < layers; i++ { for i := 0; i < layers; i++ {
masks = append(masks, mask) masks = append(masks, mask)
} }
maskManager := finalmask.NewUdpmaskManager(masks)
dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) { client, err := net.ListenPacket("udp", "127.0.0.1:0")
udpAddr, err := net.ResolveUDPAddr("udp", dest.NetAddr())
if err != nil {
return nil, nil, err
}
conn, err := gonet.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
return nil, nil, err
}
return conn, udpAddr, nil
}
listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
return gonet.ListenPacket(addr.Network(), addr.String())
}
finalMask := finalmask.NewFinalMask(nil, masks, nil, nil, dialUDP, listenPacket)
server, err := finalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: net.LocalHostIP.IP()})
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
t.Cleanup(func() { server.Close() })
clientConn, err := finalMask.DialUDP(context.Background(), net.UDPDestination(net.IPAddress(server.LocalAddr().(*net.UDPAddr).IP), net.Port(server.LocalAddr().(*net.UDPAddr).Port))) client, err = maskManager.WrapPacketConnClient(client)
if err != nil {
t.Fatal(err)
}
server, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
server, err = maskManager.WrapPacketConnServer(server)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
t.Cleanup(func() { clientConn.Close() })
client := clientConn.(*finalmask.PacketConnWrapper).PacketConn
_ = client.SetDeadline(time.Now().Add(time.Second)) _ = client.SetDeadline(time.Now().Add(time.Second))
_ = server.SetDeadline(time.Now().Add(time.Second)) _ = server.SetDeadline(time.Now().Add(time.Second))
@@ -405,20 +397,21 @@ func TestUDPcustomStaticHeaderWireShape(t *testing.T) {
{Rand: 1, RandMin: 0x30, RandMax: 0x40}, {Rand: 1, RandMin: 0x30, RandMax: 0x40},
}, },
} }
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
defer clientRaw.Close() defer clientRaw.Close()
serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
defer serverRaw.Close() defer serverRaw.Close()
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil) client, err := maskManager.WrapPacketConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -649,11 +642,11 @@ func TestSudokuBDD(t *testing.T) {
Ascii: "prefer_ascii", Ascii: "prefer_ascii",
} }
clientRaw, serverRaw := gonet.Pipe() clientRaw, serverRaw := net.Pipe()
defer clientRaw.Close() defer clientRaw.Close()
defer serverRaw.Close() defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) clientConn, err := cfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -690,11 +683,11 @@ func TestSudokuBDD(t *testing.T) {
PaddingMax: 0, PaddingMax: 0,
} }
clientRaw, serverRaw := gonet.Pipe() clientRaw, serverRaw := net.Pipe()
defer clientRaw.Close() defer clientRaw.Close()
defer serverRaw.Close() defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) clientConn, err := cfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -745,10 +738,10 @@ func TestSudokuBDD(t *testing.T) {
countWireBytes := func(wrapServer func(net.Conn, *sudoku.Config) (net.Conn, error), cfg *sudoku.Config) int64 { countWireBytes := func(wrapServer func(net.Conn, *sudoku.Config) (net.Conn, error), cfg *sudoku.Config) int64 {
t.Helper() t.Helper()
clientRaw, serverRaw := gonet.Pipe() clientRaw, serverRaw := net.Pipe()
watchedServerRaw := &countingConn{Conn: serverRaw} watchedServerRaw := &countingConn{Conn: serverRaw}
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) clientConn, err := cfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -800,11 +793,11 @@ func TestSudokuBDD(t *testing.T) {
CustomTables: []string{"xpxvvpvv", "vxpvxvvp"}, CustomTables: []string{"xpxvvpvv", "vxpvxvvp"},
} }
clientRaw, serverRaw := gonet.Pipe() clientRaw, serverRaw := net.Pipe()
defer clientRaw.Close() defer clientRaw.Close()
defer serverRaw.Close() defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) clientConn, err := cfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -842,11 +835,11 @@ func TestSudokuBDD(t *testing.T) {
PaddingMax: 0, PaddingMax: 0,
} }
clientRaw, serverRaw := gonet.Pipe() clientRaw, serverRaw := net.Pipe()
defer clientRaw.Close() defer clientRaw.Close()
defer serverRaw.Close() defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) clientConn, err := cfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -875,6 +868,19 @@ func TestSudokuBDD(t *testing.T) {
} }
}) })
t.Run("GivenSudokuUDPMask_WhenNotInnermost_ThenWrapFails", func(t *testing.T) {
cfg := &sudoku.Config{Password: "sudoku-udp"}
raw, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer raw.Close()
if _, err := cfg.WrapPacketConnClient(raw, 0, 1); err == nil {
t.Fatal("expected innermost check failure")
}
})
t.Run("GivenSudokuMultiTableUDPMask_WhenClientSendsMultipleDatagrams_ThenPayloadMatches", func(t *testing.T) { t.Run("GivenSudokuMultiTableUDPMask_WhenClientSendsMultipleDatagrams_ThenPayloadMatches", func(t *testing.T) {
cfg := &sudoku.Config{ cfg := &sudoku.Config{
Password: "sudoku-udp-multi", Password: "sudoku-udp-multi",
@@ -883,24 +889,25 @@ func TestSudokuBDD(t *testing.T) {
PaddingMin: 0, PaddingMin: 0,
PaddingMax: 0, PaddingMax: 0,
} }
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
defer clientRaw.Close() defer clientRaw.Close()
serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0") serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
defer serverRaw.Close() defer serverRaw.Close()
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil) client, err := maskManager.WrapPacketConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil) server, err := maskManager.WrapPacketConnServer(serverRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -954,7 +961,7 @@ func TestSudokuBDD(t *testing.T) {
} }
defer serverRaw.Close() defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) clientConn, err := cfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -1001,11 +1008,11 @@ func TestSudokuBDD(t *testing.T) {
Ascii: "prefer_entropy", Ascii: "prefer_entropy",
} }
clientRaw, serverRaw := gonet.Pipe() clientRaw, serverRaw := net.Pipe()
defer clientRaw.Close() defer clientRaw.Close()
defer serverRaw.Close() defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) clientConn, err := cfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -1025,11 +1032,11 @@ func TestSudokuBDD(t *testing.T) {
Ascii: "prefer_entropy", Ascii: "prefer_entropy",
} }
clientRaw, serverRaw := gonet.Pipe() clientRaw, serverRaw := net.Pipe()
defer clientRaw.Close() defer clientRaw.Close()
defer serverRaw.Close() defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil) clientConn, err := cfg.WrapConnClient(clientRaw)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
+10 -7
View File
@@ -1,17 +1,20 @@
package udphop package udphop
import ( import (
"net"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/finalmask"
) )
func (c *Config) HandleDial() {} func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
_, ok1 := raw.(*internet.FakePacketConn)
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { if level != 0 || ok1 {
return NewUDPHopConn(c, dest, dialer) return nil, errors.New("udphop requires being at the outermost level")
}
return NewUDPHopConn(c, raw)
} }
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return nil, errors.New("udphop: client only") return nil, errors.New("udphop: client only")
} }
@@ -7,6 +7,7 @@
package udphop package udphop
import ( import (
internet "github.com/xtls/xray-core/transport/internet"
protoreflect "google.golang.org/protobuf/reflect/protoreflect" protoreflect "google.golang.org/protobuf/reflect/protoreflect"
protoimpl "google.golang.org/protobuf/runtime/protoimpl" protoimpl "google.golang.org/protobuf/runtime/protoimpl"
reflect "reflect" reflect "reflect"
@@ -23,13 +24,14 @@ const (
type Config struct { type Config struct {
state protoimpl.MessageState `protogen:"open.v1"` state protoimpl.MessageState `protogen:"open.v1"`
Sockopt *internet.SocketConfig `protobuf:"bytes,1,opt,name=sockopt,proto3" json:"sockopt,omitempty"`
Local bool `protobuf:"varint,2,opt,name=local,proto3" json:"local,omitempty"` Local bool `protobuf:"varint,2,opt,name=local,proto3" json:"local,omitempty"`
Remote bool `protobuf:"varint,3,opt,name=remote,proto3" json:"remote,omitempty"` Remote bool `protobuf:"varint,3,opt,name=remote,proto3" json:"remote,omitempty"`
RemoteOnce bool `protobuf:"varint,4,opt,name=remote_once,json=remoteOnce,proto3" json:"remote_once,omitempty"` RemoteOnce bool `protobuf:"varint,4,opt,name=remote_once,json=remoteOnce,proto3" json:"remote_once,omitempty"`
IntervalMin int64 `protobuf:"varint,5,opt,name=interval_min,json=intervalMin,proto3" json:"interval_min,omitempty"` IntervalMin int64 `protobuf:"varint,5,opt,name=interval_min,json=intervalMin,proto3" json:"interval_min,omitempty"`
IntervalMax int64 `protobuf:"varint,6,opt,name=interval_max,json=intervalMax,proto3" json:"interval_max,omitempty"` IntervalMax int64 `protobuf:"varint,6,opt,name=interval_max,json=intervalMax,proto3" json:"interval_max,omitempty"`
RemoteIPs []string `protobuf:"bytes,7,rep,name=remoteIPs,proto3" json:"remoteIPs,omitempty"` RemotePorts []uint32 `protobuf:"varint,7,rep,packed,name=remote_ports,json=remotePorts,proto3" json:"remote_ports,omitempty"`
RemotePorts []uint32 `protobuf:"varint,8,rep,packed,name=remote_ports,json=remotePorts,proto3" json:"remote_ports,omitempty"` RemoteIPs []string `protobuf:"bytes,8,rep,name=remoteIPs,proto3" json:"remoteIPs,omitempty"`
unknownFields protoimpl.UnknownFields unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache sizeCache protoimpl.SizeCache
} }
@@ -64,6 +66,13 @@ func (*Config) Descriptor() ([]byte, []int) {
return file_transport_internet_finalmask_udphop_config_proto_rawDescGZIP(), []int{0} return file_transport_internet_finalmask_udphop_config_proto_rawDescGZIP(), []int{0}
} }
func (x *Config) GetSockopt() *internet.SocketConfig {
if x != nil {
return x.Sockopt
}
return nil
}
func (x *Config) GetLocal() bool { func (x *Config) GetLocal() bool {
if x != nil { if x != nil {
return x.Local return x.Local
@@ -99,16 +108,16 @@ func (x *Config) GetIntervalMax() int64 {
return 0 return 0
} }
func (x *Config) GetRemoteIPs() []string { func (x *Config) GetRemotePorts() []uint32 {
if x != nil { if x != nil {
return x.RemoteIPs return x.RemotePorts
} }
return nil return nil
} }
func (x *Config) GetRemotePorts() []uint32 { func (x *Config) GetRemoteIPs() []string {
if x != nil { if x != nil {
return x.RemotePorts return x.RemoteIPs
} }
return nil return nil
} }
@@ -117,16 +126,17 @@ var File_transport_internet_finalmask_udphop_config_proto protoreflect.FileDescr
const file_transport_internet_finalmask_udphop_config_proto_rawDesc = "" + const file_transport_internet_finalmask_udphop_config_proto_rawDesc = "" +
"\n" + "\n" +
"0transport/internet/finalmask/udphop/config.proto\x12(xray.transport.internet.finalmask.udphop\"\xe4\x01\n" + "0transport/internet/finalmask/udphop/config.proto\x12(xray.transport.internet.finalmask.udphop\x1a\x1ftransport/internet/config.proto\"\x9f\x02\n" +
"\x06Config\x12\x14\n" + "\x06Config\x12?\n" +
"\asockopt\x18\x01 \x01(\v2%.xray.transport.internet.SocketConfigR\asockopt\x12\x14\n" +
"\x05local\x18\x02 \x01(\bR\x05local\x12\x16\n" + "\x05local\x18\x02 \x01(\bR\x05local\x12\x16\n" +
"\x06remote\x18\x03 \x01(\bR\x06remote\x12\x1f\n" + "\x06remote\x18\x03 \x01(\bR\x06remote\x12\x1f\n" +
"\vremote_once\x18\x04 \x01(\bR\n" + "\vremote_once\x18\x04 \x01(\bR\n" +
"remoteOnce\x12!\n" + "remoteOnce\x12!\n" +
"\finterval_min\x18\x05 \x01(\x03R\vintervalMin\x12!\n" + "\finterval_min\x18\x05 \x01(\x03R\vintervalMin\x12!\n" +
"\finterval_max\x18\x06 \x01(\x03R\vintervalMax\x12\x1c\n" + "\finterval_max\x18\x06 \x01(\x03R\vintervalMax\x12!\n" +
"\tremoteIPs\x18\a \x03(\tR\tremoteIPs\x12!\n" + "\fremote_ports\x18\a \x03(\rR\vremotePorts\x12\x1c\n" +
"\fremote_ports\x18\b \x03(\rR\vremotePortsJ\x04\b\x01\x10\x02B\x9a\x01\n" + "\tremoteIPs\x18\b \x03(\tR\tremoteIPsB\x9a\x01\n" +
",com.xray.transport.internet.finalmask.udphopP\x01Z=github.com/xtls/xray-core/transport/internet/finalmask/udphop\xaa\x02(Xray.Transport.Internet.Finalmask.Udphopb\x06proto3" ",com.xray.transport.internet.finalmask.udphopP\x01Z=github.com/xtls/xray-core/transport/internet/finalmask/udphop\xaa\x02(Xray.Transport.Internet.Finalmask.Udphopb\x06proto3"
var ( var (
@@ -143,14 +153,16 @@ func file_transport_internet_finalmask_udphop_config_proto_rawDescGZIP() []byte
var file_transport_internet_finalmask_udphop_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1) var file_transport_internet_finalmask_udphop_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
var file_transport_internet_finalmask_udphop_config_proto_goTypes = []any{ var file_transport_internet_finalmask_udphop_config_proto_goTypes = []any{
(*Config)(nil), // 0: xray.transport.internet.finalmask.udphop.Config (*Config)(nil), // 0: xray.transport.internet.finalmask.udphop.Config
(*internet.SocketConfig)(nil), // 1: xray.transport.internet.SocketConfig
} }
var file_transport_internet_finalmask_udphop_config_proto_depIdxs = []int32{ var file_transport_internet_finalmask_udphop_config_proto_depIdxs = []int32{
0, // [0:0] is the sub-list for method output_type 1, // 0: xray.transport.internet.finalmask.udphop.Config.sockopt:type_name -> xray.transport.internet.SocketConfig
0, // [0:0] is the sub-list for method input_type 1, // [1:1] is the sub-list for method output_type
0, // [0:0] is the sub-list for extension type_name 1, // [1:1] is the sub-list for method input_type
0, // [0:0] is the sub-list for extension extendee 1, // [1:1] is the sub-list for extension type_name
0, // [0:0] is the sub-list for field type_name 1, // [1:1] is the sub-list for extension extendee
0, // [0:1] is the sub-list for field type_name
} }
func init() { file_transport_internet_finalmask_udphop_config_proto_init() } func init() { file_transport_internet_finalmask_udphop_config_proto_init() }
@@ -6,14 +6,16 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/udph
option java_package = "com.xray.transport.internet.finalmask.udphop"; option java_package = "com.xray.transport.internet.finalmask.udphop";
option java_multiple_files = true; option java_multiple_files = true;
import "transport/internet/config.proto";
message Config { message Config {
reserved 1; xray.transport.internet.SocketConfig sockopt = 1;
bool local = 2; bool local = 2;
bool remote = 3; bool remote = 3;
bool remote_once = 4; bool remote_once = 4;
int64 interval_min = 5; int64 interval_min = 5;
int64 interval_max = 6; int64 interval_max = 6;
repeated string remoteIPs = 7; repeated uint32 remote_ports = 7;
repeated uint32 remote_ports = 8; repeated string remoteIPs = 8;
} }
+97 -81
View File
@@ -6,7 +6,9 @@ import (
goerrors "errors" goerrors "errors"
"io" "io"
mrand "math/rand" mrand "math/rand"
gonet "net"
"net/netip" "net/netip"
"reflect"
"sync" "sync"
"time" "time"
@@ -14,6 +16,8 @@ import (
"github.com/xtls/xray-core/common/crypto" "github.com/xtls/xray-core/common/crypto"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/net/cnc"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet/finalmask"
) )
@@ -30,14 +34,16 @@ type packet struct {
} }
type udpHopConn struct { type udpHopConn struct {
dialer *finalmask.Dialer conn net.PacketConn
local bool sockopt *internet.SocketConfig
remote bool local bool
remote bool
remoteOnce bool
intervalMin int64 intervalMin int64
intervalMax int64 intervalMax int64
remoteIPs []netip.Prefix
remotePorts []uint32 remotePorts []uint32
remoteIPs []netip.Prefix
deadline time.Time deadline time.Time
readDeadline time.Time readDeadline time.Time
@@ -49,10 +55,10 @@ type udpHopConn struct {
readCh chan packet readCh chan packet
closeCh chan struct{} closeCh chan struct{}
wg sync.WaitGroup wg sync.WaitGroup
mu sync.RWMutex mu sync.Mutex
} }
func NewUDPHopConn(c *Config, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func NewUDPHopConn(c *Config, raw net.PacketConn) (net.PacketConn, error) {
if c.IntervalMin < 5 || c.IntervalMax < 5 { if c.IntervalMin < 5 || c.IntervalMax < 5 {
return nil, errors.New("invalid interval") return nil, errors.New("invalid interval")
} }
@@ -60,40 +66,22 @@ func NewUDPHopConn(c *Config, dest *net.Destination, dialer *finalmask.Dialer) (
for _, ip := range c.RemoteIPs { for _, ip := range c.RemoteIPs {
remoteIPs = append(remoteIPs, netip.MustParsePrefix(ip)) remoteIPs = append(remoteIPs, netip.MustParsePrefix(ip))
} }
remotePorts := c.RemotePorts conn := &udpHopConn{
if c.Remote || c.RemoteOnce { conn: raw,
if len(remoteIPs) > 0 { sockopt: c.Sockopt,
dest.Address = net.IPAddress(randPrefix(remoteIPs[mrand.Intn(len(remoteIPs))])) local: c.Local,
} remote: c.Remote,
if len(remotePorts) > 0 { remoteOnce: c.RemoteOnce,
dest.Port = net.Port(remotePorts[mrand.Intn(len(remotePorts))])
}
}
conn, err := dialer.DialUDP(*dest)
if err != nil {
return nil, err
}
cur := conn.(*finalmask.PacketConnWrapper).PacketConn
addr := conn.RemoteAddr().(*net.UDPAddr)
client := &udpHopConn{
dialer: dialer,
local: c.Local,
remote: c.Remote,
intervalMin: c.IntervalMin, intervalMin: c.IntervalMin,
intervalMax: c.IntervalMax, intervalMax: c.IntervalMax,
remotePorts: c.RemotePorts,
remoteIPs: remoteIPs, remoteIPs: remoteIPs,
remotePorts: remotePorts,
cur: cur,
addr: addr,
readCh: make(chan packet), readCh: make(chan packet),
closeCh: make(chan struct{}), closeCh: make(chan struct{}),
} }
go client.run() return conn, nil
client.wg.Add(1)
go client.recv(client.cur)
return client, nil
} }
func (c *udpHopConn) closed() bool { func (c *udpHopConn) closed() bool {
@@ -105,67 +93,61 @@ func (c *udpHopConn) closed() bool {
} }
} }
func (c *udpHopConn) run() { func (c *udpHopConn) hop(addr *net.UDPAddr) {
ticker := time.NewTicker(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
defer ticker.Stop()
for {
select {
case <-c.closeCh:
return
case <-ticker.C:
ticker.Reset(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
c.hop()
}
}
}
func (c *udpHopConn) hop() {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed() { if c.closed() {
return return
} }
oldIP := c.addr.IP newAddr := &net.UDPAddr{IP: addr.IP, Port: addr.Port}
oldPort := c.addr.Port newConn := c.conn
if c.remote { if c.remote || c.remoteOnce && c.addr == nil {
if len(c.remoteIPs) > 0 {
c.addr.IP = randPrefix(c.remoteIPs[mrand.Intn(len(c.remoteIPs))])
}
if len(c.remotePorts) > 0 { if len(c.remotePorts) > 0 {
c.addr.Port = int(c.remotePorts[mrand.Intn(len(c.remotePorts))]) newAddr.Port = int(c.remotePorts[mrand.Intn(len(c.remotePorts))])
}
if len(c.remoteIPs) > 0 {
newAddr.IP = randPrefix(c.remoteIPs[mrand.Intn(len(c.remoteIPs))])
} }
} }
if c.local { if c.local {
conn, err := c.dialer.DialUDP(net.UDPDestination(net.IPAddress(c.addr.IP), net.Port(c.addr.Port))) raw, err := internet.DialSystem(context.Background(), net.UDPDestination(net.IPAddress(newAddr.IP), net.Port(newAddr.Port)), c.sockopt)
if err != nil { if err != nil {
c.addr.IP = oldIP
c.addr.Port = oldPort
errors.LogErrorInner(context.Background(), err, "hop err") errors.LogErrorInner(context.Background(), err, "hop err")
return return
} }
conn.SetDeadline(c.deadline) switch c := raw.(type) {
conn.SetReadDeadline(c.readDeadline) case *internet.PacketConnWrapper:
conn.SetWriteDeadline(c.writeDeadline) newConn = c.PacketConn
case *cnc.Connection:
newConn = &internet.FakePacketConn{Conn: c}
default:
panic(reflect.TypeOf(c))
}
newConn.SetDeadline(c.deadline)
newConn.SetReadDeadline(c.readDeadline)
newConn.SetWriteDeadline(c.writeDeadline)
if c.pre != nil { if c.pre != nil {
_ = c.pre.Close() _ = c.pre.Close()
} }
c.pre = c.cur c.pre = c.cur
c.cur = conn.(*finalmask.PacketConnWrapper).PacketConn
c.wg.Add(1) c.wg.Add(1)
go c.recv(c.cur) go c.recv(newConn)
} }
c.addr = newAddr
c.cur = newConn
} }
func (c *udpHopConn) recv(conn net.PacketConn) { func (c *udpHopConn) recv(conn net.PacketConn) {
defer c.wg.Done() defer c.wg.Done()
for { for {
if c.closed() {
return
}
p := pool.Get().([]byte) p := pool.Get().([]byte)
n, addr, err := conn.ReadFrom(p) n, addr, err := conn.ReadFrom(p)
if err != nil { if err != nil {
pool.Put(p[:cap(p)]) pool.Put(p[:cap(p)])
if c.closed() { if goerrors.Is(err, io.EOF) || goerrors.Is(err, io.ErrClosedPipe) || goerrors.Is(err, gonet.ErrClosed) {
return break
} }
var netErr net.Error var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() { if goerrors.As(err, &netErr) && netErr.Timeout() {
@@ -174,10 +156,9 @@ func (c *udpHopConn) recv(conn net.PacketConn) {
case <-c.closeCh: case <-c.closeCh:
return return
} }
continue
} }
errors.LogErrorInner(context.Background(), err, "recv err") errors.LogErrorInner(context.Background(), err, "recv err")
return continue
} }
select { select {
case c.readCh <- packet{p: p[:n], addr: addr}: case c.readCh <- packet{p: p[:n], addr: addr}:
@@ -188,6 +169,22 @@ func (c *udpHopConn) recv(conn net.PacketConn) {
} }
} }
func (c *udpHopConn) hopLoop() {
ticker := time.NewTicker(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
defer ticker.Stop()
for {
select {
case <-ticker.C:
ticker.Reset(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
c.mu.Lock()
c.hop(c.addr)
c.mu.Unlock()
case <-c.closeCh:
return
}
}
}
func (c *udpHopConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) { func (c *udpHopConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
packet, ok := <-c.readCh packet, ok := <-c.readCh
if ok { if ok {
@@ -197,12 +194,21 @@ func (c *udpHopConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
} }
return n, packet.addr, packet.err return n, packet.addr, packet.err
} }
return 0, nil, io.ErrClosedPipe return 0, nil, io.EOF
} }
func (c *udpHopConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { func (c *udpHopConn) WriteTo(p []byte, addr net.Addr) (n int, err error) {
c.mu.RLock() c.mu.Lock()
defer c.mu.RUnlock() defer c.mu.Unlock()
if c.cur == nil {
c.hop(addr.(*net.UDPAddr))
if c.cur == nil {
return 0, nil
}
go c.hopLoop()
}
_, err = c.cur.WriteTo(p, c.addr) _, err = c.cur.WriteTo(p, c.addr)
if err != nil { if err != nil {
errors.LogErrorInner(context.Background(), err, "send err") errors.LogErrorInner(context.Background(), err, "send err")
@@ -221,12 +227,15 @@ func (c *udpHopConn) Close() error {
if c.pre != nil { if c.pre != nil {
_ = c.pre.Close() _ = c.pre.Close()
} }
_ = c.cur.Close() if c.cur != nil {
_ = c.cur.Close()
}
_ = c.conn.Close()
c.wg.Wait() c.wg.Wait()
select { select {
case packet := <-c.readCh: case p := <-c.readCh:
if packet.p != nil { if p.p != nil {
pool.Put(packet.p[:cap(packet.p)]) pool.Put(p.p[:cap(p.p)])
} }
default: default:
} }
@@ -235,9 +244,7 @@ func (c *udpHopConn) Close() error {
} }
func (c *udpHopConn) LocalAddr() net.Addr { func (c *udpHopConn) LocalAddr() net.Addr {
c.mu.RLock() return c.conn.LocalAddr()
defer c.mu.RUnlock()
return c.cur.LocalAddr()
} }
func (c *udpHopConn) SetDeadline(t time.Time) error { func (c *udpHopConn) SetDeadline(t time.Time) error {
@@ -247,7 +254,10 @@ func (c *udpHopConn) SetDeadline(t time.Time) error {
if c.pre != nil { if c.pre != nil {
_ = c.pre.SetDeadline(t) _ = c.pre.SetDeadline(t)
} }
return c.cur.SetDeadline(t) if c.cur != nil {
_ = c.cur.SetDeadline(t)
}
return nil
} }
func (c *udpHopConn) SetReadDeadline(t time.Time) error { func (c *udpHopConn) SetReadDeadline(t time.Time) error {
@@ -257,7 +267,10 @@ func (c *udpHopConn) SetReadDeadline(t time.Time) error {
if c.pre != nil { if c.pre != nil {
_ = c.pre.SetReadDeadline(t) _ = c.pre.SetReadDeadline(t)
} }
return c.cur.SetReadDeadline(t) if c.cur != nil {
_ = c.cur.SetReadDeadline(t)
}
return nil
} }
func (c *udpHopConn) SetWriteDeadline(t time.Time) error { func (c *udpHopConn) SetWriteDeadline(t time.Time) error {
@@ -267,7 +280,10 @@ func (c *udpHopConn) SetWriteDeadline(t time.Time) error {
if c.pre != nil { if c.pre != nil {
_ = c.pre.SetWriteDeadline(t) _ = c.pre.SetWriteDeadline(t)
} }
return c.cur.SetWriteDeadline(t) if c.cur != nil {
_ = c.cur.SetWriteDeadline(t)
}
return nil
} }
func randPrefix(p netip.Prefix) []byte { func randPrefix(p netip.Prefix) []byte {
+13 -6
View File
@@ -1,14 +1,21 @@
package xdns package xdns
import ( import (
"github.com/xtls/xray-core/common/net" "net"
"github.com/xtls/xray-core/transport/internet/finalmask"
) )
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClient(c, conn) // _, ok1 := raw.(*internet.FakePacketConn)
// _, ok2 := raw.(*udphop.UdpHopPacketConn)
// if level != 0 || ok1 || ok2 {
// return nil, errors.New("xdns requires being at the outermost level")
// }
return NewConnClient(c, raw)
} }
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServer(c, conn) // if level != 0 {
// return nil, errors.New("xdns requires being at the outermost level")
// }
return NewConnServer(c, raw)
} }
+34 -28
View File
@@ -8,7 +8,8 @@ import (
goerrors "errors" goerrors "errors"
"fmt" "fmt"
"io" "io"
mrand "math/rand" mathrand "math/rand"
"net"
"net/netip" "net/netip"
"sync" "sync"
"time" "time"
@@ -16,7 +17,6 @@ import (
"github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet/finalmask"
"golang.org/x/net/icmp" "golang.org/x/net/icmp"
"golang.org/x/net/ipv4" "golang.org/x/net/ipv4"
@@ -36,11 +36,11 @@ type packet struct {
} }
type xicmpConnClient struct { type xicmpConnClient struct {
conn net.PacketConn
icmp4 *icmp.PacketConn icmp4 *icmp.PacketConn
icmp6 *icmp.PacketConn icmp6 *icmp.PacketConn
udp bool udp bool
ips []netip.Addr ips []netip.Addr
ip net.IP
clientID [8]byte clientID [8]byte
id int id int
seq int seq int
@@ -50,7 +50,7 @@ type xicmpConnClient struct {
mu sync.Mutex mu sync.Mutex
} }
func NewConnClient(c *Config, dest *net.Destination) (net.PacketConn, error) { func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
var icmp4, icmp6 *icmp.PacketConn var icmp4, icmp6 *icmp.PacketConn
var err4, err6 error var err4, err6 error
if c.DGRAM { if c.DGRAM {
@@ -69,24 +69,17 @@ func NewConnClient(c *Config, dest *net.Destination) (net.PacketConn, error) {
ips = append(ips, netip.MustParseAddr(ip)) ips = append(ips, netip.MustParseAddr(ip))
} }
var ip net.IP
if len(ips) > 0 {
ip = ips[mrand.Intn(len(ips))].AsSlice()
} else {
ip = dest.Address.IP()
}
var clientID [8]byte var clientID [8]byte
common.Must2(rand.Read(clientID[:])) common.Must2(rand.Read(clientID[:]))
conn := &xicmpConnClient{ conn := &xicmpConnClient{
conn: raw,
icmp4: icmp4, icmp4: icmp4,
icmp6: icmp6, icmp6: icmp6,
udp: c.DGRAM, udp: c.DGRAM,
ips: ips, ips: ips,
ip: ip,
clientID: clientID, clientID: clientID,
id: mrand.Intn(65536), id: mathrand.Intn(65536),
seq: 1, seq: 1,
readCh: make(chan packet), readCh: make(chan packet),
closeCh: make(chan struct{}), closeCh: make(chan struct{}),
@@ -99,6 +92,10 @@ func NewConnClient(c *Config, dest *net.Destination) (net.PacketConn, error) {
return conn, nil return conn, nil
} }
func (c *xicmpConnClient) ring(a, b uint16) uint16 {
return min(a-b, b-a)
}
func (c *xicmpConnClient) closed() bool { func (c *xicmpConnClient) closed() bool {
select { select {
case <-c.closeCh: case <-c.closeCh:
@@ -113,11 +110,12 @@ func (c *xicmpConnClient) recv4() {
var b [finalmask.UDPSize]byte var b [finalmask.UDPSize]byte
for { for {
if c.closed() {
return
}
n, addr, err := c.icmp4.ReadFrom(b[:]) n, addr, err := c.icmp4.ReadFrom(b[:])
if err != nil { if err != nil {
if c.closed() {
return
}
var netErr net.Error var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() { if goerrors.As(err, &netErr) && netErr.Timeout() {
select { select {
@@ -127,10 +125,9 @@ func (c *xicmpConnClient) recv4() {
case <-c.closeCh: case <-c.closeCh:
return return
} }
continue
} }
errors.LogErrorInner(context.Background(), err, "recv err 4") errors.LogErrorInner(context.Background(), err, "recv4 err")
return continue
} }
msg, err := icmp.ParseMessage(1, b[:n]) msg, err := icmp.ParseMessage(1, b[:n])
@@ -153,6 +150,10 @@ func (c *xicmpConnClient) recv4() {
continue continue
} }
if c.ring(uint16(echo.Seq), uint16(c.seq)) > 1000 {
continue
}
if len(echo.Data) > 8 && bytes.Equal(echo.Data[:8], c.clientID[:]) { if len(echo.Data) > 8 && bytes.Equal(echo.Data[:8], c.clientID[:]) {
continue continue
} }
@@ -181,11 +182,12 @@ func (c *xicmpConnClient) recv6() {
var b [finalmask.UDPSize]byte var b [finalmask.UDPSize]byte
for { for {
if c.closed() {
return
}
n, addr, err := c.icmp6.ReadFrom(b[:]) n, addr, err := c.icmp6.ReadFrom(b[:])
if err != nil { if err != nil {
if c.closed() {
return
}
var netErr net.Error var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() { if goerrors.As(err, &netErr) && netErr.Timeout() {
select { select {
@@ -195,10 +197,9 @@ func (c *xicmpConnClient) recv6() {
case <-c.closeCh: case <-c.closeCh:
return return
} }
continue
} }
errors.LogErrorInner(context.Background(), err, "recv err 6") errors.LogErrorInner(context.Background(), err, "recv6 err")
return continue
} }
msg, err := icmp.ParseMessage(58, b[:n]) msg, err := icmp.ParseMessage(58, b[:n])
@@ -221,6 +222,10 @@ func (c *xicmpConnClient) recv6() {
continue continue
} }
if c.ring(uint16(echo.Seq), uint16(c.seq)) > 1000 {
continue
}
if len(echo.Data) > 8 && bytes.Equal(echo.Data[:8], c.clientID[:]) { if len(echo.Data) > 8 && bytes.Equal(echo.Data[:8], c.clientID[:]) {
continue continue
} }
@@ -268,9 +273,9 @@ func (c *xicmpConnClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
c.seq %= 65536 c.seq %= 65536
c.mu.Unlock() c.mu.Unlock()
ip := c.ip ip := addr.(*net.UDPAddr).IP
if len(c.ips) > 0 { if len(c.ips) > 0 {
ip = c.ips[mrand.Intn(len(c.ips))].AsSlice() ip = c.ips[mathrand.Intn(len(c.ips))].AsSlice()
} }
if c.udp { if c.udp {
@@ -309,6 +314,7 @@ func (c *xicmpConnClient) Close() error {
close(c.closeCh) close(c.closeCh)
_ = c.icmp4.Close() _ = c.icmp4.Close()
_ = c.icmp6.Close() _ = c.icmp6.Close()
_ = c.conn.Close()
c.wg.Wait() c.wg.Wait()
select { select {
case p := <-c.readCh: case p := <-c.readCh:
@@ -322,7 +328,7 @@ func (c *xicmpConnClient) Close() error {
} }
func (c *xicmpConnClient) LocalAddr() net.Addr { func (c *xicmpConnClient) LocalAddr() net.Addr {
return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} return c.conn.LocalAddr()
} }
func (c *xicmpConnClient) SetDeadline(t time.Time) error { func (c *xicmpConnClient) SetDeadline(t time.Time) error {
+13 -13
View File
@@ -1,23 +1,23 @@
package xicmp package xicmp
import ( import (
"errors" "net"
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet"
) )
func (c *Config) HandleDial() {} func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
_, ok1 := raw.(*internet.FakePacketConn)
func (c *Config) HandleListen() {} if level != 0 || ok1 {
return nil, errors.New("xicmp requires being at the outermost level")
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
if dest.Address.Family().IsDomain() && len(c.IPs) == 0 {
return nil, errors.New("empty ip addresses")
} }
return NewConnClient(c, dest) return NewConnClient(c, raw)
} }
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServer(c) if level != 0 {
return nil, errors.New("xicmp requires being at the outermost level")
}
return NewConnServer(c, raw)
} }
+17 -14
View File
@@ -37,6 +37,7 @@ type record struct {
} }
type xicmpConnServer struct { type xicmpConnServer struct {
conn net.PacketConn
icmp4 *icmp.PacketConn icmp4 *icmp.PacketConn
icmp6 *icmp.PacketConn icmp6 *icmp.PacketConn
ips map[netip.Addr]struct{} ips map[netip.Addr]struct{}
@@ -47,7 +48,7 @@ type xicmpConnServer struct {
mu sync.Mutex mu sync.Mutex
} }
func NewConnServer(c *Config) (net.PacketConn, error) { func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
icmp4, err := icmp.ListenPacket("ip4:icmp", "0.0.0.0") icmp4, err := icmp.ListenPacket("ip4:icmp", "0.0.0.0")
if err != nil { if err != nil {
return nil, err return nil, err
@@ -63,6 +64,7 @@ func NewConnServer(c *Config) (net.PacketConn, error) {
} }
conn := &xicmpConnServer{ conn := &xicmpConnServer{
conn: raw,
icmp4: icmp4, icmp4: icmp4,
icmp6: icmp6, icmp6: icmp6,
ips: ips, ips: ips,
@@ -113,11 +115,12 @@ func (c *xicmpConnServer) recv4() {
var b [finalmask.UDPSize]byte var b [finalmask.UDPSize]byte
for { for {
if c.closed() {
return
}
n, addr, err := c.icmp4.ReadFrom(b[:]) n, addr, err := c.icmp4.ReadFrom(b[:])
if err != nil { if err != nil {
if c.closed() {
return
}
var netErr net.Error var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() { if goerrors.As(err, &netErr) && netErr.Timeout() {
select { select {
@@ -127,10 +130,9 @@ func (c *xicmpConnServer) recv4() {
case <-c.closeCh: case <-c.closeCh:
return return
} }
continue
} }
errors.LogErrorInner(context.Background(), err, "recv err 4") errors.LogErrorInner(context.Background(), err, "recv4 err")
return continue
} }
msg, err := icmp.ParseMessage(1, b[:n]) msg, err := icmp.ParseMessage(1, b[:n])
@@ -193,11 +195,12 @@ func (c *xicmpConnServer) recv6() {
var b [finalmask.UDPSize]byte var b [finalmask.UDPSize]byte
for { for {
if c.closed() {
return
}
n, addr, err := c.icmp6.ReadFrom(b[:]) n, addr, err := c.icmp6.ReadFrom(b[:])
if err != nil { if err != nil {
if c.closed() {
return
}
var netErr net.Error var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() { if goerrors.As(err, &netErr) && netErr.Timeout() {
select { select {
@@ -207,10 +210,9 @@ func (c *xicmpConnServer) recv6() {
case <-c.closeCh: case <-c.closeCh:
return return
} }
continue
} }
errors.LogErrorInner(context.Background(), err, "recv err 6") errors.LogErrorInner(context.Background(), err, "recv6 err")
return continue
} }
msg, err := icmp.ParseMessage(58, b[:n]) msg, err := icmp.ParseMessage(58, b[:n])
@@ -328,6 +330,7 @@ func (c *xicmpConnServer) Close() error {
close(c.closeCh) close(c.closeCh)
_ = c.icmp4.Close() _ = c.icmp4.Close()
_ = c.icmp6.Close() _ = c.icmp6.Close()
_ = c.conn.Close()
c.wg.Wait() c.wg.Wait()
select { select {
case p := <-c.readCh: case p := <-c.readCh:
@@ -341,7 +344,7 @@ func (c *xicmpConnServer) Close() error {
} }
func (c *xicmpConnServer) LocalAddr() net.Addr { func (c *xicmpConnServer) LocalAddr() net.Addr {
return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} return c.conn.LocalAddr()
} }
func (c *xicmpConnServer) SetDeadline(t time.Time) error { func (c *xicmpConnServer) SetDeadline(t time.Time) error {
@@ -39,6 +39,7 @@ type record struct {
} }
type xicmpConnServer struct { type xicmpConnServer struct {
conn net.PacketConn
icmp4 *icmp.PacketConn icmp4 *icmp.PacketConn
icmp6 *icmp.PacketConn icmp6 *icmp.PacketConn
ipv4PC *ipv4.PacketConn ipv4PC *ipv4.PacketConn
@@ -51,7 +52,7 @@ type xicmpConnServer struct {
mu sync.Mutex mu sync.Mutex
} }
func NewConnServer(c *Config) (net.PacketConn, error) { func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
icmp4, err := icmp.ListenPacket("ip4:icmp", "0.0.0.0") icmp4, err := icmp.ListenPacket("ip4:icmp", "0.0.0.0")
if err != nil { if err != nil {
return nil, err return nil, err
@@ -67,6 +68,7 @@ func NewConnServer(c *Config) (net.PacketConn, error) {
} }
conn := &xicmpConnServer{ conn := &xicmpConnServer{
conn: raw,
icmp4: icmp4, icmp4: icmp4,
icmp6: icmp6, icmp6: icmp6,
ipv4PC: icmp4.IPv4PacketConn(), ipv4PC: icmp4.IPv4PacketConn(),
@@ -122,11 +124,12 @@ func (c *xicmpConnServer) recv4() {
var b [finalmask.UDPSize]byte var b [finalmask.UDPSize]byte
for { for {
if c.closed() {
return
}
n, cm, addr, err := c.ipv4PC.ReadFrom(b[:]) n, cm, addr, err := c.ipv4PC.ReadFrom(b[:])
if err != nil { if err != nil {
if c.closed() {
return
}
var netErr net.Error var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() { if goerrors.As(err, &netErr) && netErr.Timeout() {
select { select {
@@ -136,10 +139,9 @@ func (c *xicmpConnServer) recv4() {
case <-c.closeCh: case <-c.closeCh:
return return
} }
continue
} }
errors.LogErrorInner(context.Background(), err, "recv err 4") errors.LogErrorInner(context.Background(), err, "recv4 err")
return continue
} }
msg, err := icmp.ParseMessage(1, b[:n]) msg, err := icmp.ParseMessage(1, b[:n])
@@ -203,11 +205,12 @@ func (c *xicmpConnServer) recv6() {
var b [finalmask.UDPSize]byte var b [finalmask.UDPSize]byte
for { for {
if c.closed() {
return
}
n, cm, addr, err := c.ipv6PC.ReadFrom(b[:]) n, cm, addr, err := c.ipv6PC.ReadFrom(b[:])
if err != nil { if err != nil {
if c.closed() {
return
}
var netErr net.Error var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() { if goerrors.As(err, &netErr) && netErr.Timeout() {
select { select {
@@ -217,10 +220,9 @@ func (c *xicmpConnServer) recv6() {
case <-c.closeCh: case <-c.closeCh:
return return
} }
continue
} }
errors.LogErrorInner(context.Background(), err, "recv err 6") errors.LogErrorInner(context.Background(), err, "recv6 err")
return continue
} }
msg, err := icmp.ParseMessage(58, b[:n]) msg, err := icmp.ParseMessage(58, b[:n])
@@ -339,6 +341,7 @@ func (c *xicmpConnServer) Close() error {
close(c.closeCh) close(c.closeCh)
_ = c.icmp4.Close() _ = c.icmp4.Close()
_ = c.icmp6.Close() _ = c.icmp6.Close()
_ = c.conn.Close()
c.wg.Wait() c.wg.Wait()
select { select {
case p := <-c.readCh: case p := <-c.readCh:
@@ -352,7 +355,7 @@ func (c *xicmpConnServer) Close() error {
} }
func (c *xicmpConnServer) LocalAddr() net.Addr { func (c *xicmpConnServer) LocalAddr() net.Addr {
return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} return c.conn.LocalAddr()
} }
func (c *xicmpConnServer) SetDeadline(t time.Time) error { func (c *xicmpConnServer) SetDeadline(t time.Time) error {
+2 -4
View File
@@ -2,12 +2,10 @@ package xmc
import ( import (
"fmt" "fmt"
"net"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
) )
func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) { func (c *Config) WrapConnClient(conn net.Conn) (net.Conn, error) {
profiles, err := profilesFromConfig(c.Profiles) profiles, err := profilesFromConfig(c.Profiles)
if err != nil { if err != nil {
return nil, fmt.Errorf("minecraft finalmask: %w", err) return nil, fmt.Errorf("minecraft finalmask: %w", err)
+11 -6
View File
@@ -83,6 +83,7 @@ func getGrpcClient(ctx context.Context, dest net.Destination, streamSettings *in
} }
tlsConfig := tls.ConfigFromStreamSettings(streamSettings) tlsConfig := tls.ConfigFromStreamSettings(streamSettings)
realityConfig := reality.ConfigFromStreamSettings(streamSettings) realityConfig := reality.ConfigFromStreamSettings(streamSettings)
sockopt := streamSettings.SocketSettings
grpcSettings := streamSettings.ProtocolSettings.(*Config) grpcSettings := streamSettings.ProtocolSettings.(*Config)
if client, found := globalDialerMap[dialerConf{dest, streamSettings}]; found && client.GetState() != connectivity.Shutdown { if client, found := globalDialerMap[dialerConf{dest, streamSettings}]; found && client.GetState() != connectivity.Shutdown {
@@ -123,13 +124,17 @@ func getGrpcClient(ctx context.Context, dest net.Destination, streamSettings *in
gctx = session.ContextWithOutbounds(gctx, session.OutboundsFromContext(ctx)) gctx = session.ContextWithOutbounds(gctx, session.OutboundsFromContext(ctx))
gctx = session.ContextWithTimeoutOnly(gctx, true) gctx = session.ContextWithTimeoutOnly(gctx, true)
var c net.Conn c, err := internet.DialSystem(gctx, net.TCPDestination(address, port), sockopt)
if streamSettings.FinalMask != nil {
c, err = streamSettings.FinalMask.DialTCP(gctx, net.TCPDestination(address, port))
} else {
c, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
}
if err == nil { if err == nil {
if streamSettings.TcpmaskManager != nil {
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(c)
if err != nil {
c.Close()
return nil, errors.New("mask err").Base(err)
}
c = newConn
}
if tlsConfig != nil { if tlsConfig != nil {
config := tlsConfig.GetTLSConfig(tls.WithDestination(dest)) config := tlsConfig.GetTLSConfig(tls.WithDestination(dest))
if fingerprint := tls.GetFingerprint(tlsConfig.Fingerprint); fingerprint != nil { if fingerprint := tls.GetFingerprint(tlsConfig.Fingerprint); fingerprint != nil {
+19 -11
View File
@@ -104,20 +104,28 @@ func Listen(ctx context.Context, address net.Address, port net.Port, settings *i
go func() { go func() {
var streamListener net.Listener var streamListener net.Listener
var err error var err error
var addr net.Addr
if port == net.Port(0) { // unix if port == net.Port(0) { // unix
addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"} streamListener, err = internet.ListenSystem(ctx, &net.UnixAddr{
Name: address.Domain(),
Net: "unix",
}, settings.SocketSettings)
if err != nil {
errors.LogErrorInner(ctx, err, "failed to listen on ", address)
return
}
} else { // tcp } else { // tcp
addr = &net.TCPAddr{IP: address.IP(), Port: int(port)} streamListener, err = internet.ListenSystem(ctx, &net.TCPAddr{
IP: address.IP(),
Port: int(port),
}, settings.SocketSettings)
if err != nil {
errors.LogErrorInner(ctx, err, "failed to listen on ", address, ":", port)
return
}
} }
if settings.FinalMask != nil {
streamListener, err = settings.FinalMask.Listen(ctx, addr) if settings.TcpmaskManager != nil {
} else { streamListener, _ = settings.TcpmaskManager.WrapListener(streamListener)
streamListener, err = internet.ListenSystem(ctx, addr, settings.SocketSettings)
}
if err != nil {
errors.LogErrorInner(ctx, err, "failed to listen on ", address, ":", port)
return
} }
errors.LogDebug(ctx, "gRPC listen for service name `"+grpcSettings.getServiceName()+"` tun `"+grpcSettings.getTunStreamName()+"` multi tun `"+grpcSettings.getTunMultiStreamName()+"`") errors.LogDebug(ctx, "gRPC listen for service name `"+grpcSettings.getServiceName()+"` tun `"+grpcSettings.getTunStreamName()+"` multi tun `"+grpcSettings.getTunMultiStreamName()+"`")
+10 -7
View File
@@ -46,18 +46,21 @@ func (c *ConnRF) Read(b []byte) (int, error) {
func dialhttpUpgrade(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (net.Conn, error) { func dialhttpUpgrade(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (net.Conn, error) {
transportConfiguration := streamSettings.ProtocolSettings.(*Config) transportConfiguration := streamSettings.ProtocolSettings.(*Config)
var pconn net.Conn pconn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
var err error
if streamSettings.FinalMask != nil {
pconn, err = streamSettings.FinalMask.DialTCP(ctx, dest)
} else {
pconn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
}
if err != nil { if err != nil {
errors.LogErrorInner(ctx, err, "failed to dial to ", dest) errors.LogErrorInner(ctx, err, "failed to dial to ", dest)
return nil, err return nil, err
} }
if streamSettings.TcpmaskManager != nil {
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(pconn)
if err != nil {
pconn.Close()
return nil, errors.New("mask err").Base(err)
}
pconn = newConn
}
var conn net.Conn var conn net.Conn
var requestURL url.URL var requestURL url.URL
tConfig := tls.ConfigFromStreamSettings(streamSettings) tConfig := tls.ConfigFromStreamSettings(streamSettings)
+19 -11
View File
@@ -124,21 +124,29 @@ func ListenHTTPUpgrade(ctx context.Context, address net.Address, port net.Port,
} }
var listener net.Listener var listener net.Listener
var err error var err error
var addr net.Addr
if port == net.Port(0) { // unix if port == net.Port(0) { // unix
addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"} listener, err = internet.ListenSystem(ctx, &net.UnixAddr{
Name: address.Domain(),
Net: "unix",
}, streamSettings.SocketSettings)
if err != nil {
return nil, errors.New("failed to listen unix domain socket(for HttpUpgrade) on ", address).Base(err)
}
errors.LogInfo(ctx, "listening unix domain socket(for HttpUpgrade) on ", address)
} else { // tcp } else { // tcp
addr = &net.TCPAddr{IP: address.IP(), Port: int(port)} listener, err = internet.ListenSystem(ctx, &net.TCPAddr{
IP: address.IP(),
Port: int(port),
}, streamSettings.SocketSettings)
if err != nil {
return nil, errors.New("failed to listen TCP(for HttpUpgrade) on ", address, ":", port).Base(err)
}
errors.LogInfo(ctx, "listening TCP(for HttpUpgrade) on ", address, ":", port)
} }
if streamSettings.FinalMask != nil {
listener, err = streamSettings.FinalMask.Listen(ctx, addr) if streamSettings.TcpmaskManager != nil {
} else { listener, _ = streamSettings.TcpmaskManager.WrapListener(listener)
listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings)
} }
if err != nil {
return nil, errors.New("failed to listen ", addr.Network(), "(for HttpUpgrade) on ", address, ":", port).Base(err)
}
errors.LogInfo(ctx, "listening ", addr.Network(), "(for HttpUpgrade) on ", address, ":", port)
if streamSettings.SocketSettings != nil && streamSettings.SocketSettings.AcceptProxyProtocol { if streamSettings.SocketSettings != nil && streamSettings.SocketSettings.AcceptProxyProtocol {
errors.LogWarning(ctx, "accepting PROXY protocol") errors.LogWarning(ctx, "accepting PROXY protocol")
+36 -35
View File
@@ -2,7 +2,7 @@ package hysteria
import ( import (
"context" "context"
gotls "crypto/tls" go_tls "crypto/tls"
"net/http" "net/http"
"net/url" "net/url"
"reflect" "reflect"
@@ -28,12 +28,12 @@ import (
type client struct { type client struct {
sync.Mutex sync.Mutex
dest net.Destination dest net.Destination
config *Config config *Config
tlsConfig *gotls.Config tlsConfig *go_tls.Config
socketConfig *internet.SocketConfig socketConfig *internet.SocketConfig
finalMask *finalmask.FinalMask udpmaskManager *finalmask.UdpmaskManager
quicParams *internet.QuicParams quicParams *internet.QuicParams
conn *quic.Conn conn *quic.Conn
tr *quic.Transport tr *quic.Transport
@@ -113,29 +113,30 @@ func (c *client) dial(ctx context.Context) error {
// } // }
var pktConn net.PacketConn var pktConn net.PacketConn
var udpAddr net.Addr var udpAddr *net.UDPAddr
if c.finalMask != nil {
conn, err := c.finalMask.DialUDP(ctx, c.dest) raw, err := internet.DialSystem(ctx, c.dest, c.socketConfig)
if err != nil {
return errors.New("failed to dial to dest").Base(err)
}
switch c := raw.(type) {
case *internet.PacketConnWrapper:
pktConn = c.PacketConn
udpAddr = raw.RemoteAddr().(*net.UDPAddr)
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port}
default:
panic(reflect.TypeOf(c))
}
if c.udpmaskManager != nil {
newConn, err := c.udpmaskManager.WrapPacketConnClient(pktConn)
if err != nil { if err != nil {
return errors.New("failed to dial to dest").Base(err) pktConn.Close()
} return errors.New("mask err").Base(err)
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
udpAddr = conn.RemoteAddr()
} else {
conn, err := internet.DialSystem(ctx, c.dest, c.socketConfig)
if err != nil {
return errors.New("failed to dial to dest").Base(err)
}
switch c := conn.(type) {
case *internet.PacketConnWrapper:
pktConn = c.PacketConn
udpAddr = c.RemoteAddr()
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
default:
panic(reflect.TypeOf(c))
} }
pktConn = newConn
} }
tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO} tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO}
@@ -149,7 +150,7 @@ func (c *client) dial(ctx context.Context) error {
rt := &http3.Transport{ rt := &http3.Transport{
TLSClientConfig: c.tlsConfig, TLSClientConfig: c.tlsConfig,
QUICConfig: quicConfig, QUICConfig: quicConfig,
Dial: func(ctx context.Context, _ string, tlsCfg *gotls.Config, cfg *quic.Config) (*quic.Conn, error) { Dial: func(ctx context.Context, _ string, tlsCfg *go_tls.Config, cfg *quic.Config) (*quic.Conn, error) {
qc, err := tr.DialEarly(ctx, udpAddr, tlsCfg, cfg) qc, err := tr.DialEarly(ctx, udpAddr, tlsCfg, cfg)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -315,12 +316,12 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
c = manager.m[dialerConf{dest, streamSettings}] c = manager.m[dialerConf{dest, streamSettings}]
if c == nil { if c == nil {
c = &client{ c = &client{
dest: dest, dest: dest,
config: streamSettings.ProtocolSettings.(*Config), config: streamSettings.ProtocolSettings.(*Config),
tlsConfig: tlsConfig.GetTLSConfig(tls.WithDestination(dest)), tlsConfig: tlsConfig.GetTLSConfig(tls.WithDestination(dest)),
socketConfig: streamSettings.SocketSettings, socketConfig: streamSettings.SocketSettings,
finalMask: streamSettings.FinalMask, udpmaskManager: streamSettings.UdpmaskManager,
quicParams: streamSettings.QuicParams, quicParams: streamSettings.QuicParams,
} }
manager.m[dialerConf{dest, streamSettings}] = c manager.m[dialerConf{dest, streamSettings}] = c
} }
+10 -7
View File
@@ -316,17 +316,20 @@ func Listen(ctx context.Context, address net.Address, port net.Port, streamSetti
quicConfig.MaxIncomingStreams = 1024 quicConfig.MaxIncomingStreams = 1024
} }
var pktConn net.PacketConn pktConn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings)
var err error
if streamSettings.FinalMask != nil {
pktConn, err = streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)})
} else {
pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings)
}
if err != nil { if err != nil {
return nil, err return nil, err
} }
if streamSettings.UdpmaskManager != nil {
newConn, err := streamSettings.UdpmaskManager.WrapPacketConnServer(pktConn)
if err != nil {
pktConn.Close()
return nil, errors.New("mask err").Base(err)
}
pktConn = newConn
}
var k *quic.StatelessResetKey var k *quic.StatelessResetKey
if !quicParams.DisableStatelessReset { if !quicParams.DisableStatelessReset {
k = &quic.StatelessResetKey{} k = &quic.StatelessResetKey{}
@@ -1,165 +0,0 @@
package hysteria
import (
"context"
"crypto/tls"
"crypto/x509"
"errors"
"net"
"runtime"
"testing"
"time"
"github.com/apernet/quic-go"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/protocol/tls/cert"
)
func TestDatagram(t *testing.T) {
run := func() (addr net.Addr, recv chan int64, cancel func()) {
cert, _ := cert.MustGenerate(nil)
Certificate := [][]byte{cert.Certificate}
PrivateKey := common.Must2(x509.ParsePKCS8PrivateKey(cert.PrivateKey))
tlsConf := &tls.Config{
Certificates: []tls.Certificate{
{
Certificate: Certificate,
PrivateKey: PrivateKey,
},
},
NextProtos: []string{"h3"},
}
quicConf := &quic.Config{
InitialStreamReceiveWindow: 8388608,
MaxStreamReceiveWindow: 8388608,
InitialConnectionReceiveWindow: 8388608 * 5 / 2,
MaxConnectionReceiveWindow: 8388608 * 5 / 2,
MaxIdleTimeout: 30 * time.Second,
MaxIncomingStreams: 1024,
DisablePathMTUDiscovery: runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin",
EnableDatagrams: true,
MaxDatagramFrameSize: MaxDatagramFrameSize,
AssumePeerMaxDatagramFrameSize: MaxDatagramFrameSize,
DisablePathManager: true,
}
pktConn := common.Must2(net.ListenPacket("udp", "127.0.0.1:0"))
tr := &quic.Transport{Conn: pktConn}
l := common.Must2(tr.Listen(tlsConf, quicConf))
recv = make(chan int64)
ctx, cancel := context.WithCancel(context.Background())
go func() {
defer pktConn.Close()
defer tr.Close()
defer l.Close()
defer close(recv)
var buf [1500]byte
for {
conn, err := l.Accept(ctx)
if err != nil {
if !errors.Is(err, context.Canceled) {
t.Error(err)
}
break
}
err = conn.SendDatagram(buf[:])
var qErr *quic.DatagramTooLargeError
if !errors.As(err, &qErr) {
t.Error(err)
}
recv <- qErr.MaxDatagramPayloadSize
defer conn.CloseWithError(0, "")
}
}()
return l.Addr(), recv, cancel
}
addr, recv, cancel := run()
t.Run("With ChromeParrot", func(t *testing.T) {
tlsConf := &tls.Config{
InsecureSkipVerify: true,
}
quicConf := &quic.Config{
InitialStreamReceiveWindow: 8388608,
MaxStreamReceiveWindow: 8388608,
InitialConnectionReceiveWindow: 8388608 * 5 / 2,
MaxConnectionReceiveWindow: 8388608 * 5 / 2,
MaxIdleTimeout: 30 * time.Second,
KeepAlivePeriod: 10 * time.Second,
DisablePathMTUDiscovery: runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin",
ChromeParrot: true,
EnableDatagrams: true,
MaxDatagramFrameSize: MaxDatagramFrameSize,
OmitMaxDatagramFrameSize: true,
DisablePathManager: true,
}
pktConn := common.Must2(net.ListenPacket("udp", "127.0.0.1:0"))
tr := &quic.Transport{Conn: pktConn, ConnectionIDGenerator: quic.ZeroLengthConnectionIDGenerator{}}
conn := common.Must2(tr.DialEarly(context.Background(), addr, tlsConf, quicConf))
defer pktConn.Close()
defer tr.Close()
defer conn.CloseWithError(0, "")
var buf [1500]byte
err := conn.SendDatagram(buf[:])
var qErr *quic.DatagramTooLargeError
if !errors.As(err, &qErr) || qErr.MaxDatagramPayloadSize != 1197 {
t.Error(err)
}
if server := <-recv; server != 1243 {
t.Error(server)
}
})
t.Run("Without ChromeParrot", func(t *testing.T) {
tlsConf := &tls.Config{
InsecureSkipVerify: true,
NextProtos: []string{"h3"},
}
quicConf := &quic.Config{
InitialStreamReceiveWindow: 8388608,
MaxStreamReceiveWindow: 8388608,
InitialConnectionReceiveWindow: 8388608 * 5 / 2,
MaxConnectionReceiveWindow: 8388608 * 5 / 2,
MaxIdleTimeout: 30 * time.Second,
KeepAlivePeriod: 10 * time.Second,
DisablePathMTUDiscovery: runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin",
ChromeParrot: false,
EnableDatagrams: true,
MaxDatagramFrameSize: MaxDatagramFrameSize,
OmitMaxDatagramFrameSize: true,
DisablePathManager: true,
}
pktConn := common.Must2(net.ListenPacket("udp", "127.0.0.1:0"))
tr := &quic.Transport{Conn: pktConn}
conn := common.Must2(tr.DialEarly(context.Background(), addr, tlsConf, quicConf))
defer pktConn.Close()
defer tr.Close()
defer conn.CloseWithError(0, "")
var buf [1500]byte
err := conn.SendDatagram(buf[:])
var qErr *quic.DatagramTooLargeError
if !errors.As(err, &qErr) || qErr.MaxDatagramPayloadSize != 1197 {
t.Error(err)
}
if server := <-recv; server != 1197 {
t.Error(server)
}
})
cancel()
}
+29 -8
View File
@@ -3,6 +3,7 @@ package kcp
import ( import (
"context" "context"
"io" "io"
reflect "reflect"
"sync/atomic" "sync/atomic"
"github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common"
@@ -10,6 +11,7 @@ import (
"github.com/xtls/xray-core/common/dice" "github.com/xtls/xray-core/common/dice"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/net/cnc"
"github.com/xtls/xray-core/transport/internet" "github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/stat" "github.com/xtls/xray-core/transport/internet/stat"
"github.com/xtls/xray-core/transport/internet/tls" "github.com/xtls/xray-core/transport/internet/tls"
@@ -49,15 +51,34 @@ func DialKCP(ctx context.Context, dest net.Destination, streamSettings *internet
dest.Network = net.Network_UDP dest.Network = net.Network_UDP
errors.LogInfo(ctx, "dialing mKCP to ", dest) errors.LogInfo(ctx, "dialing mKCP to ", dest)
var conn net.Conn conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
var err error
if streamSettings.FinalMask != nil {
conn, err = streamSettings.FinalMask.DialUDP(ctx, dest)
} else {
conn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
}
if err != nil { if err != nil {
return nil, errors.New("failed to dial to dest: ", err).AtWarning().Base(err) return nil, errors.New("failed to dial to dest: ", err).Base(err)
}
if streamSettings.UdpmaskManager != nil {
var pktConn net.PacketConn
var udpAddr *net.UDPAddr
switch c := conn.(type) {
case *internet.PacketConnWrapper:
pktConn = c.PacketConn
udpAddr = c.RemoteAddr().(*net.UDPAddr)
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port}
default:
panic(reflect.TypeOf(c))
}
newConn, err := streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn)
if err != nil {
pktConn.Close()
return nil, errors.New("mask err").Base(err)
}
pktConn = newConn
conn = &internet.PacketConnWrapper{
PacketConn: pktConn,
Dest: udpAddr,
}
} }
kcpSettings := streamSettings.ProtocolSettings.(*Config) kcpSettings := streamSettings.ProtocolSettings.(*Config)
+23 -47
View File
@@ -1,12 +1,7 @@
package internet package internet
import ( import (
"context"
"reflect"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/net/cnc"
"github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet/finalmask"
) )
@@ -17,7 +12,8 @@ type MemoryStreamConfig struct {
ProtocolSettings interface{} ProtocolSettings interface{}
SecurityType string SecurityType string
SecuritySettings interface{} SecuritySettings interface{}
FinalMask *finalmask.FinalMask TcpmaskManager *finalmask.TcpmaskManager
UdpmaskManager *finalmask.UdpmaskManager
QuicParams *QuicParams QuicParams *QuicParams
SocketSettings *SocketConfig SocketSettings *SocketConfig
DownloadSettings *MemoryStreamConfig DownloadSettings *MemoryStreamConfig
@@ -55,53 +51,33 @@ func ToMemoryStreamConfig(s *StreamConfig) (*MemoryStreamConfig, error) {
mss.SecuritySettings = ess mss.SecuritySettings = ess
} }
var tcpMasks []finalmask.TCPMask if s != nil && len(s.Tcpmasks) > 0 {
var udpMasks []finalmask.UDPMask var masks []finalmask.Tcpmask
for _, msg := range s.Tcpmasks {
if s != nil { instance, err := msg.GetInstance()
for i := range s.Tcpmasks { if err != nil {
instance := common.Must2(s.Tcpmasks[i].GetInstance()) return nil, err
tcpMasks = append(tcpMasks, instance.(finalmask.TCPMask)) }
} masks = append(masks, instance.(finalmask.Tcpmask))
for i := range s.Udpmasks {
instance := common.Must2(s.Udpmasks[i].GetInstance())
udpMasks = append(udpMasks, instance.(finalmask.UDPMask))
} }
mss.TcpmaskManager = finalmask.NewTcpmaskManager(masks)
} }
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
return DialSystem(ctx, dest, mss.SocketSettings)
}
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
return ListenSystem(ctx, addr, mss.SocketSettings)
}
dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) {
conn, err := DialSystem(ctx, dest, mss.SocketSettings)
if err != nil {
return nil, nil, err
}
var newConn net.PacketConn
var udpAddr net.Addr
switch c := conn.(type) {
case *PacketConnWrapper:
newConn = c.PacketConn
udpAddr = conn.RemoteAddr()
case *cnc.Connection:
newConn = &FakePacketConn{Conn: c}
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}, Port: 0}
default:
panic(reflect.TypeOf(c))
}
return newConn, udpAddr, nil
}
listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
return ListenSystemPacket(ctx, addr, mss.SocketSettings)
}
mss.FinalMask = finalmask.NewFinalMask(tcpMasks, udpMasks, dialTCP, listen, dialUDP, listenPacket)
if s != nil && s.QuicParams != nil { if s != nil && s.QuicParams != nil {
mss.QuicParams = s.QuicParams mss.QuicParams = s.QuicParams
} }
if s != nil && len(s.Udpmasks) > 0 {
var masks []finalmask.Udpmask
for _, msg := range s.Udpmasks {
instance, err := msg.GetInstance()
if err != nil {
return nil, err
}
masks = append(masks, instance.(finalmask.Udpmask))
}
mss.UdpmaskManager = finalmask.NewUdpmaskManager(masks)
}
return mss, nil return mss, nil
} }

Some files were not shown because too many files have changed in this diff Show More