mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-02 20:08:12 +03:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2cf8eba4cc | ||
|
|
bf3230ec51 |
@@ -2,33 +2,43 @@ package outbound
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"slices"
|
||||||
"sort"
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/app/proxyman"
|
"github.com/xtls/xray-core/app/proxyman"
|
||||||
"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/utils"
|
||||||
"github.com/xtls/xray-core/core"
|
"github.com/xtls/xray-core/core"
|
||||||
"github.com/xtls/xray-core/features/outbound"
|
"github.com/xtls/xray-core/features/outbound"
|
||||||
|
"github.com/xtls/xray-core/features/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Manager is to manage all outbound handlers.
|
// Manager is to manage all outbound handlers.
|
||||||
type Manager struct {
|
type Manager struct {
|
||||||
access sync.RWMutex
|
defaultHandler atomic.Pointer[outbound.Handler]
|
||||||
defaultHandler outbound.Handler
|
taggedHandler *utils.TypedSyncMap[string, outbound.Handler]
|
||||||
taggedHandler map[string]outbound.Handler
|
untaggedHandlers atomic.Pointer[[]outbound.Handler]
|
||||||
untaggedHandlers []outbound.Handler
|
running atomic.Bool
|
||||||
running bool
|
tagsCache *utils.TypedSyncMap[string, []string]
|
||||||
tagsCache *sync.Map
|
balancerPicker routing.BalancerPicker
|
||||||
}
|
}
|
||||||
|
|
||||||
// New creates a new Manager.
|
// New creates a new Manager.
|
||||||
func New(ctx context.Context, config *proxyman.OutboundConfig) (*Manager, error) {
|
func New(ctx context.Context, config *proxyman.OutboundConfig) (*Manager, error) {
|
||||||
m := &Manager{
|
m := &Manager{
|
||||||
taggedHandler: make(map[string]outbound.Handler),
|
taggedHandler: utils.NewTypedSyncMap[string, outbound.Handler](),
|
||||||
tagsCache: &sync.Map{},
|
|
||||||
}
|
}
|
||||||
|
m.tagsCache = utils.NewTypedSyncMap[string, []string]()
|
||||||
|
empty := make([]outbound.Handler, 0)
|
||||||
|
m.untaggedHandlers.Store(&empty)
|
||||||
|
_ = core.OptionalFeatures(ctx, func(router routing.Router) {
|
||||||
|
if picker, ok := router.(routing.BalancerPicker); ok {
|
||||||
|
m.balancerPicker = picker
|
||||||
|
}
|
||||||
|
})
|
||||||
return m, nil
|
return m, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -39,20 +49,25 @@ func (m *Manager) Type() interface{} {
|
|||||||
|
|
||||||
// Start implements core.Feature
|
// Start implements core.Feature
|
||||||
func (m *Manager) Start() error {
|
func (m *Manager) Start() error {
|
||||||
m.access.Lock()
|
m.running.Store(true)
|
||||||
defer m.access.Unlock()
|
|
||||||
|
|
||||||
m.running = true
|
var startErr error
|
||||||
|
m.taggedHandler.Range(func(_ string, h outbound.Handler) bool {
|
||||||
for _, h := range m.taggedHandler {
|
|
||||||
if err := h.Start(); err != nil {
|
if err := h.Start(); err != nil {
|
||||||
return err
|
startErr = err
|
||||||
|
return false
|
||||||
}
|
}
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
if startErr != nil {
|
||||||
|
return startErr
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, h := range m.untaggedHandlers {
|
if untagged := m.untaggedHandlers.Load(); untagged != nil {
|
||||||
if err := h.Start(); err != nil {
|
for _, h := range *untagged {
|
||||||
return err
|
if err := h.Start(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -61,18 +76,18 @@ func (m *Manager) Start() error {
|
|||||||
|
|
||||||
// Close implements core.Feature
|
// Close implements core.Feature
|
||||||
func (m *Manager) Close() error {
|
func (m *Manager) Close() error {
|
||||||
m.access.Lock()
|
m.running.Store(false)
|
||||||
defer m.access.Unlock()
|
|
||||||
|
|
||||||
m.running = false
|
|
||||||
|
|
||||||
var errs []error
|
var errs []error
|
||||||
for _, h := range m.taggedHandler {
|
m.taggedHandler.Range(func(_ string, h outbound.Handler) bool {
|
||||||
errs = append(errs, h.Close())
|
errs = append(errs, h.Close())
|
||||||
}
|
return true
|
||||||
|
})
|
||||||
|
|
||||||
for _, h := range m.untaggedHandlers {
|
if untagged := m.untaggedHandlers.Load(); untagged != nil {
|
||||||
errs = append(errs, h.Close())
|
for _, h := range *untagged {
|
||||||
|
errs = append(errs, h.Close())
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return errors.Combine(errs...)
|
return errors.Combine(errs...)
|
||||||
@@ -80,47 +95,69 @@ func (m *Manager) Close() error {
|
|||||||
|
|
||||||
// GetDefaultHandler implements outbound.Manager.
|
// GetDefaultHandler implements outbound.Manager.
|
||||||
func (m *Manager) GetDefaultHandler() outbound.Handler {
|
func (m *Manager) GetDefaultHandler() outbound.Handler {
|
||||||
m.access.RLock()
|
if h := m.defaultHandler.Load(); h != nil {
|
||||||
defer m.access.RUnlock()
|
return *h
|
||||||
|
|
||||||
if m.defaultHandler == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return m.defaultHandler
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetHandler implements outbound.Manager.
|
|
||||||
func (m *Manager) GetHandler(tag string) outbound.Handler {
|
|
||||||
m.access.RLock()
|
|
||||||
defer m.access.RUnlock()
|
|
||||||
if handler, found := m.taggedHandler[tag]; found {
|
|
||||||
return handler
|
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetHandler implements outbound.Manager.
|
||||||
|
func (m *Manager) GetHandler(tag string) outbound.Handler {
|
||||||
|
if handler, found := m.taggedHandler.Load(tag); found {
|
||||||
|
return handler
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(tag, "balancer:") && m.balancerPicker != nil {
|
||||||
|
targetTag := m.tryGetOutboundTagWithBalancer(tag, nil)
|
||||||
|
if targetTag == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if handler, found := m.taggedHandler.Load(targetTag); found {
|
||||||
|
return handler
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *Manager) tryGetOutboundTagWithBalancer(tag string, parents []string) string {
|
||||||
|
balancerTag := tag[len("balancer:"):]
|
||||||
|
targetTag, err := m.balancerPicker.GetBalancerOutboundTag(balancerTag)
|
||||||
|
if err != nil {
|
||||||
|
errors.LogWarning(context.Background(), "failed to pick outbound from balancer [", balancerTag, "]: ", err)
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(targetTag, "balancer:") {
|
||||||
|
if slices.Contains(parents, balancerTag) {
|
||||||
|
errors.LogWarning(context.Background(), "detected balancer loop for [", balancerTag, "]")
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return m.tryGetOutboundTagWithBalancer(targetTag, append(parents, balancerTag))
|
||||||
|
}
|
||||||
|
return targetTag
|
||||||
|
}
|
||||||
|
|
||||||
// AddHandler implements outbound.Manager.
|
// AddHandler implements outbound.Manager.
|
||||||
func (m *Manager) AddHandler(ctx context.Context, handler outbound.Handler) error {
|
func (m *Manager) AddHandler(ctx context.Context, handler outbound.Handler) error {
|
||||||
m.access.Lock()
|
m.defaultHandler.CompareAndSwap(nil, &handler)
|
||||||
defer m.access.Unlock()
|
|
||||||
|
|
||||||
m.tagsCache = &sync.Map{}
|
|
||||||
|
|
||||||
if m.defaultHandler == nil {
|
|
||||||
m.defaultHandler = handler
|
|
||||||
}
|
|
||||||
|
|
||||||
tag := handler.Tag()
|
tag := handler.Tag()
|
||||||
if len(tag) > 0 {
|
if len(tag) > 0 {
|
||||||
if _, found := m.taggedHandler[tag]; found {
|
if _, found := m.taggedHandler.LoadOrStore(tag, handler); found {
|
||||||
return errors.New("existing tag found: " + tag)
|
return errors.New("existing tag found: " + tag)
|
||||||
}
|
}
|
||||||
m.taggedHandler[tag] = handler
|
|
||||||
} else {
|
} else {
|
||||||
m.untaggedHandlers = append(m.untaggedHandlers, handler)
|
for {
|
||||||
|
oldUntagged := m.untaggedHandlers.Load()
|
||||||
|
newUntagged := make([]outbound.Handler, 0, len(*oldUntagged)+1)
|
||||||
|
newUntagged = append(newUntagged, *oldUntagged...)
|
||||||
|
newUntagged = append(newUntagged, handler)
|
||||||
|
if m.untaggedHandlers.CompareAndSwap(oldUntagged, &newUntagged) {
|
||||||
|
m.tagsCache.Clear()
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if m.running {
|
if m.running.Load() {
|
||||||
return handler.Start()
|
return handler.Start()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -132,54 +169,49 @@ func (m *Manager) RemoveHandler(ctx context.Context, tag string) error {
|
|||||||
if tag == "" {
|
if tag == "" {
|
||||||
return common.ErrNoClue
|
return common.ErrNoClue
|
||||||
}
|
}
|
||||||
m.access.Lock()
|
|
||||||
defer m.access.Unlock()
|
|
||||||
|
|
||||||
m.tagsCache = &sync.Map{}
|
m.taggedHandler.Delete(tag)
|
||||||
|
if cur := m.defaultHandler.Load(); cur != nil && (*cur).Tag() == tag {
|
||||||
delete(m.taggedHandler, tag)
|
m.defaultHandler.CompareAndSwap(cur, nil)
|
||||||
if m.defaultHandler != nil && m.defaultHandler.Tag() == tag {
|
|
||||||
m.defaultHandler = nil
|
|
||||||
}
|
}
|
||||||
|
m.tagsCache.Clear()
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ListHandlers implements outbound.Manager.
|
// ListHandlers implements outbound.Manager.
|
||||||
func (m *Manager) ListHandlers(ctx context.Context) []outbound.Handler {
|
func (m *Manager) ListHandlers(ctx context.Context) []outbound.Handler {
|
||||||
m.access.RLock()
|
var response []outbound.Handler
|
||||||
defer m.access.RUnlock()
|
if untagged := m.untaggedHandlers.Load(); untagged != nil {
|
||||||
|
response = slices.Clone(*untagged)
|
||||||
response := make([]outbound.Handler, len(m.untaggedHandlers))
|
|
||||||
copy(response, m.untaggedHandlers)
|
|
||||||
|
|
||||||
for _, v := range m.taggedHandler {
|
|
||||||
response = append(response, v)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
m.taggedHandler.Range(func(_ string, v outbound.Handler) bool {
|
||||||
|
response = append(response, v)
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
|
||||||
return response
|
return response
|
||||||
}
|
}
|
||||||
|
|
||||||
// Select implements outbound.HandlerSelector.
|
// Select implements outbound.HandlerSelector.
|
||||||
func (m *Manager) Select(selectors []string) []string {
|
func (m *Manager) Select(selectors []string) []string {
|
||||||
key := strings.Join(selectors, ",")
|
key := strings.Join(selectors, ",")
|
||||||
if cache, ok := m.tagsCache.Load(key); ok {
|
if result, ok := m.tagsCache.Load(key); ok {
|
||||||
return cache.([]string)
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
m.access.RLock()
|
|
||||||
defer m.access.RUnlock()
|
|
||||||
|
|
||||||
tags := make([]string, 0, len(selectors))
|
tags := make([]string, 0, len(selectors))
|
||||||
|
|
||||||
for tag := range m.taggedHandler {
|
m.taggedHandler.Range(func(tag string, _ outbound.Handler) bool {
|
||||||
for _, selector := range selectors {
|
for _, selector := range selectors {
|
||||||
if strings.HasPrefix(tag, selector) {
|
if strings.HasPrefix(tag, selector) {
|
||||||
tags = append(tags, tag)
|
tags = append(tags, tag)
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
return true
|
||||||
|
})
|
||||||
|
|
||||||
sort.Strings(tags)
|
sort.Strings(tags)
|
||||||
m.tagsCache.Store(key, tags)
|
m.tagsCache.Store(key, tags)
|
||||||
|
|||||||
@@ -165,3 +165,11 @@ func (r *Router) GetOverrideTarget(tag string) (string, error) {
|
|||||||
}
|
}
|
||||||
return "", errors.New("cannot find tag")
|
return "", errors.New("cannot find tag")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// PickOutbound implements routing.BalancerPicker
|
||||||
|
func (r *Router) GetBalancerOutboundTag(balancerTag string) (string, error) {
|
||||||
|
if b, ok := (*r.balancers.Load())[balancerTag]; ok {
|
||||||
|
return b.PickOutbound()
|
||||||
|
}
|
||||||
|
return "", errors.New("cannot find tag")
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -8,3 +8,8 @@ type BalancerOverrider interface {
|
|||||||
type BalancerPrincipleTarget interface {
|
type BalancerPrincipleTarget interface {
|
||||||
GetPrincipleTarget(tag string) ([]string, error)
|
GetPrincipleTarget(tag string) ([]string, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// BalancerPicker return the outbound tag selected by the given balancer
|
||||||
|
type BalancerPicker interface {
|
||||||
|
GetBalancerOutboundTag(balancerTag string) (string, error)
|
||||||
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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=
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -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
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+4
-33
@@ -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))
|
||||||
|
|||||||
+138
-113
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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{}),
|
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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))
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,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)
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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()+"`")
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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()
|
|
||||||
}
|
|
||||||
@@ -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,17 +51,36 @@ 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).AtWarning().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)
|
||||||
|
|
||||||
reader := &KCPPacketReader{}
|
reader := &KCPPacketReader{}
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -18,7 +18,6 @@ import (
|
|||||||
"regexp"
|
"regexp"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
|
||||||
"time"
|
"time"
|
||||||
"unsafe"
|
"unsafe"
|
||||||
|
|
||||||
@@ -37,18 +36,6 @@ import (
|
|||||||
|
|
||||||
type Conn struct {
|
type Conn struct {
|
||||||
*reality.Conn
|
*reality.Conn
|
||||||
suppressCloseNotify atomic.Bool
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) SuppressCloseNotify() {
|
|
||||||
c.suppressCloseNotify.Store(true)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) Close() error {
|
|
||||||
if c.suppressCloseNotify.Load() {
|
|
||||||
return c.Conn.NetConn().Close()
|
|
||||||
}
|
|
||||||
return c.Conn.Close()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Conn) HandshakeAddress() net.Address {
|
func (c *Conn) HandshakeAddress() net.Address {
|
||||||
@@ -69,22 +56,10 @@ func Server(c net.Conn, config *reality.Config) (net.Conn, error) {
|
|||||||
|
|
||||||
type UConn struct {
|
type UConn struct {
|
||||||
*utls.UConn
|
*utls.UConn
|
||||||
Config *Config
|
Config *Config
|
||||||
ServerName string
|
ServerName string
|
||||||
AuthKey []byte
|
AuthKey []byte
|
||||||
Verified bool
|
Verified bool
|
||||||
suppressCloseNotify atomic.Bool
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UConn) SuppressCloseNotify() {
|
|
||||||
c.suppressCloseNotify.Store(true)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UConn) Close() error {
|
|
||||||
if c.suppressCloseNotify.Load() {
|
|
||||||
return c.NetConn().Close()
|
|
||||||
}
|
|
||||||
return c.UConn.Close()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *UConn) HandshakeAddress() net.Address {
|
func (c *UConn) HandshakeAddress() net.Address {
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptrace"
|
"net/http/httptrace"
|
||||||
"net/url"
|
"net/url"
|
||||||
"reflect"
|
reflect "reflect"
|
||||||
"runtime"
|
"runtime"
|
||||||
"strconv"
|
"strconv"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -25,7 +25,6 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/signal/done"
|
"github.com/xtls/xray-core/common/signal/done"
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
"github.com/xtls/xray-core/transport/internet/browser_dialer"
|
"github.com/xtls/xray-core/transport/internet/browser_dialer"
|
||||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
|
||||||
"github.com/xtls/xray-core/transport/internet/hysteria/congestion"
|
"github.com/xtls/xray-core/transport/internet/hysteria/congestion"
|
||||||
"github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr"
|
"github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr"
|
||||||
"github.com/xtls/xray-core/transport/internet/reality"
|
"github.com/xtls/xray-core/transport/internet/reality"
|
||||||
@@ -117,17 +116,20 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea
|
|||||||
transportConfig := streamSettings.ProtocolSettings.(*Config)
|
transportConfig := streamSettings.ProtocolSettings.(*Config)
|
||||||
|
|
||||||
dialContext := func(ctxInner context.Context) (net.Conn, error) {
|
dialContext := func(ctxInner context.Context) (net.Conn, error) {
|
||||||
var conn net.Conn
|
conn, err := internet.DialSystem(ctxInner, dest, streamSettings.SocketSettings)
|
||||||
var err error
|
|
||||||
if streamSettings.FinalMask != nil {
|
|
||||||
conn, err = streamSettings.FinalMask.DialTCP(ctxInner, dest)
|
|
||||||
} else {
|
|
||||||
conn, err = internet.DialSystem(ctxInner, dest, streamSettings.SocketSettings)
|
|
||||||
}
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if streamSettings.TcpmaskManager != nil {
|
||||||
|
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(conn)
|
||||||
|
if err != nil {
|
||||||
|
conn.Close()
|
||||||
|
return nil, errors.New("mask err").Base(err)
|
||||||
|
}
|
||||||
|
conn = newConn
|
||||||
|
}
|
||||||
|
|
||||||
if realityConfig != nil {
|
if realityConfig != nil {
|
||||||
return reality.UClient(conn, realityConfig, ctxInner, dest)
|
return reality.UClient(conn, realityConfig, ctxInner, dest)
|
||||||
}
|
}
|
||||||
@@ -194,29 +196,30 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea
|
|||||||
TLSClientConfig: gotlsConfig,
|
TLSClientConfig: gotlsConfig,
|
||||||
Dial: func(ctx context.Context, addr string, tlsCfg *gotls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
Dial: func(ctx context.Context, addr string, tlsCfg *gotls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||||
var pktConn net.PacketConn
|
var pktConn net.PacketConn
|
||||||
var udpAddr net.Addr
|
var udpAddr *net.UDPAddr
|
||||||
if streamSettings.FinalMask != nil {
|
|
||||||
conn, err := streamSettings.FinalMask.DialUDP(ctx, dest)
|
raw, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||||
|
if err != nil {
|
||||||
|
return nil, 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 streamSettings.UdpmaskManager != nil {
|
||||||
|
newConn, err := 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
|
|
||||||
udpAddr = conn.RemoteAddr()
|
|
||||||
} else {
|
|
||||||
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("failed to dial to dest").Base(err)
|
|
||||||
}
|
|
||||||
switch c := conn.(type) {
|
|
||||||
case *internet.PacketConnWrapper:
|
|
||||||
pktConn = c.PacketConn
|
|
||||||
udpAddr = c.RemoteAddr()
|
|
||||||
case *cnc.Connection:
|
|
||||||
pktConn = &internet.FakePacketConn{Conn: c}
|
|
||||||
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
|
||||||
default:
|
|
||||||
panic(reflect.TypeOf(c))
|
|
||||||
}
|
}
|
||||||
|
pktConn = newConn
|
||||||
}
|
}
|
||||||
|
|
||||||
tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO}
|
tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO}
|
||||||
|
|||||||
@@ -463,17 +463,31 @@ func ListenXH(ctx context.Context, address net.Address, port net.Port, streamSet
|
|||||||
l.isH3 = len(tlsConfig.NextProtos) == 1 && tlsConfig.NextProtos[0] == "h3"
|
l.isH3 = len(tlsConfig.NextProtos) == 1 && tlsConfig.NextProtos[0] == "h3"
|
||||||
|
|
||||||
var err error
|
var err error
|
||||||
if l.isH3 {
|
if port == net.Port(0) { // unix
|
||||||
var pktConn net.PacketConn
|
l.listener, err = internet.ListenSystem(ctx, &net.UnixAddr{
|
||||||
var err error
|
Name: address.Domain(),
|
||||||
if streamSettings.FinalMask != nil {
|
Net: "unix",
|
||||||
pktConn, err = streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)})
|
}, streamSettings.SocketSettings)
|
||||||
} else {
|
if err != nil {
|
||||||
pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings)
|
return nil, errors.New("failed to listen UNIX domain socket for XHTTP on ", address).Base(err)
|
||||||
}
|
}
|
||||||
|
errors.LogInfo(ctx, "listening UNIX domain socket for XHTTP on ", address)
|
||||||
|
} else if l.isH3 { // quic
|
||||||
|
Conn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{
|
||||||
|
IP: address.IP(),
|
||||||
|
Port: int(port),
|
||||||
|
}, streamSettings.SocketSettings)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to listen UDP for XHTTP/3 on ", address, ":", port).Base(err)
|
return nil, errors.New("failed to listen UDP for XHTTP/3 on ", address, ":", port).Base(err)
|
||||||
}
|
}
|
||||||
|
if streamSettings.UdpmaskManager != nil {
|
||||||
|
newConn, err := streamSettings.UdpmaskManager.WrapPacketConnServer(Conn)
|
||||||
|
if err != nil {
|
||||||
|
Conn.Close()
|
||||||
|
return nil, errors.New("mask err").Base(err)
|
||||||
|
}
|
||||||
|
Conn = newConn
|
||||||
|
}
|
||||||
|
|
||||||
quicParams := streamSettings.QuicParams
|
quicParams := streamSettings.QuicParams
|
||||||
if quicParams == nil {
|
if quicParams == nil {
|
||||||
@@ -498,7 +512,7 @@ func ListenXH(ctx context.Context, address net.Address, port net.Port, streamSet
|
|||||||
common.Must2(rand.Read((*k)[:]))
|
common.Must2(rand.Read((*k)[:]))
|
||||||
}
|
}
|
||||||
|
|
||||||
tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO, StatelessResetKey: k}
|
tr := &quic.Transport{Conn: Conn, DisableGSO: quicParams.DisableGSO, StatelessResetKey: k}
|
||||||
|
|
||||||
l.h3listener, err = tr.ListenEarly(tlsConfig, quicConfig)
|
l.h3listener, err = tr.ListenEarly(tlsConfig, quicConfig)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -520,24 +534,21 @@ func ListenXH(ctx context.Context, address net.Address, port net.Port, streamSet
|
|||||||
errors.LogErrorInner(ctx, err, "failed to serve HTTP/3 for XHTTP/3")
|
errors.LogErrorInner(ctx, err, "failed to serve HTTP/3 for XHTTP/3")
|
||||||
}
|
}
|
||||||
_ = tr.Close()
|
_ = tr.Close()
|
||||||
_ = pktConn.Close()
|
_ = Conn.Close()
|
||||||
}()
|
}()
|
||||||
} else {
|
} else { // tcp
|
||||||
var addr net.Addr
|
l.listener, err = internet.ListenSystem(ctx, &net.TCPAddr{
|
||||||
if port == net.Port(0) { // unix
|
IP: address.IP(),
|
||||||
addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"}
|
Port: int(port),
|
||||||
} else { // tcp
|
}, streamSettings.SocketSettings)
|
||||||
addr = &net.TCPAddr{IP: address.IP(), Port: int(port)}
|
|
||||||
}
|
|
||||||
if streamSettings.FinalMask != nil {
|
|
||||||
l.listener, err = streamSettings.FinalMask.Listen(ctx, addr)
|
|
||||||
} else {
|
|
||||||
l.listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings)
|
|
||||||
}
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to listen ", addr.Network(), " for XHTTP on ", address, ":", port).Base(err)
|
return nil, errors.New("failed to listen TCP for XHTTP on ", address, ":", port).Base(err)
|
||||||
}
|
}
|
||||||
errors.LogInfo(ctx, "listening ", addr.Network(), " for XHTTP on ", address, ":", port)
|
errors.LogInfo(ctx, "listening TCP for XHTTP on ", address, ":", port)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !l.isH3 && streamSettings.TcpmaskManager != nil {
|
||||||
|
l.listener, _ = streamSettings.TcpmaskManager.WrapListener(l.listener)
|
||||||
}
|
}
|
||||||
|
|
||||||
// tcp/unix (h1/h2)
|
// tcp/unix (h1/h2)
|
||||||
|
|||||||
@@ -235,5 +235,5 @@ func (c *FakePacketConn) WriteTo(p []byte, _ net.Addr) (n int, err error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *FakePacketConn) LocalAddr() net.Addr {
|
func (c *FakePacketConn) LocalAddr() net.Addr {
|
||||||
return &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
return &net.UDPAddr{IP: c.Conn.LocalAddr().(*net.TCPAddr).IP, Port: c.Conn.LocalAddr().(*net.TCPAddr).Port}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -19,15 +19,18 @@ import (
|
|||||||
// Dial dials a new TCP connection to the given destination.
|
// Dial dials a new TCP connection to the given destination.
|
||||||
func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (stat.Connection, error) {
|
func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (stat.Connection, error) {
|
||||||
errors.LogInfo(ctx, "dialing TCP to ", dest)
|
errors.LogInfo(ctx, "dialing TCP 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.DialTCP(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").Base(err)
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if streamSettings.TcpmaskManager != nil {
|
||||||
|
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(conn)
|
||||||
|
if err != nil {
|
||||||
|
conn.Close()
|
||||||
|
return nil, errors.New("mask err").Base(err)
|
||||||
|
}
|
||||||
|
conn = newConn
|
||||||
}
|
}
|
||||||
|
|
||||||
if config := tls.ConfigFromStreamSettings(streamSettings); config != nil {
|
if config := tls.ConfigFromStreamSettings(streamSettings); config != nil {
|
||||||
|
|||||||
@@ -41,21 +41,29 @@ func ListenTCP(ctx context.Context, address net.Address, port net.Port, streamSe
|
|||||||
}
|
}
|
||||||
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{
|
||||||
} else { // tcp
|
Name: address.Domain(),
|
||||||
addr = &net.TCPAddr{IP: address.IP(), Port: int(port)}
|
Net: "unix",
|
||||||
}
|
}, streamSettings.SocketSettings)
|
||||||
if streamSettings.FinalMask != nil {
|
if err != nil {
|
||||||
listener, err = streamSettings.FinalMask.Listen(ctx, addr)
|
return nil, errors.New("failed to listen Unix Domain Socket on ", address).Base(err)
|
||||||
|
}
|
||||||
|
errors.LogInfo(ctx, "listening Unix Domain Socket on ", address)
|
||||||
} else {
|
} else {
|
||||||
listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings)
|
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 on ", address, ":", port).Base(err)
|
||||||
|
}
|
||||||
|
errors.LogInfo(ctx, "listening TCP on ", address, ":", port)
|
||||||
}
|
}
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("failed to listen ", addr.Network(), " on ", address, ":", port).Base(err)
|
if streamSettings.TcpmaskManager != nil {
|
||||||
|
listener, _ = streamSettings.TcpmaskManager.WrapListener(listener)
|
||||||
}
|
}
|
||||||
errors.LogInfo(ctx, "listening ", addr.Network(), " 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")
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ import (
|
|||||||
"crypto/tls"
|
"crypto/tls"
|
||||||
"math/big"
|
"math/big"
|
||||||
"slices"
|
"slices"
|
||||||
"sync/atomic"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
utls "github.com/refraction-networking/utls"
|
utls "github.com/refraction-networking/utls"
|
||||||
@@ -30,19 +29,11 @@ var (
|
|||||||
|
|
||||||
type Conn struct {
|
type Conn struct {
|
||||||
*tls.Conn
|
*tls.Conn
|
||||||
suppressCloseNotify atomic.Bool
|
|
||||||
}
|
}
|
||||||
|
|
||||||
const tlsCloseTimeout = 250 * time.Millisecond
|
const tlsCloseTimeout = 250 * time.Millisecond
|
||||||
|
|
||||||
func (c *Conn) SuppressCloseNotify() {
|
|
||||||
c.suppressCloseNotify.Store(true)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) Close() error {
|
func (c *Conn) Close() error {
|
||||||
if c.suppressCloseNotify.Load() {
|
|
||||||
return c.Conn.NetConn().Close()
|
|
||||||
}
|
|
||||||
timer := time.AfterFunc(tlsCloseTimeout, func() {
|
timer := time.AfterFunc(tlsCloseTimeout, func() {
|
||||||
c.Conn.NetConn().Close()
|
c.Conn.NetConn().Close()
|
||||||
})
|
})
|
||||||
@@ -83,19 +74,11 @@ func Server(c net.Conn, config *tls.Config) net.Conn {
|
|||||||
|
|
||||||
type UConn struct {
|
type UConn struct {
|
||||||
*utls.UConn
|
*utls.UConn
|
||||||
suppressCloseNotify atomic.Bool
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var _ Interface = (*UConn)(nil)
|
var _ Interface = (*UConn)(nil)
|
||||||
|
|
||||||
func (c *UConn) SuppressCloseNotify() {
|
|
||||||
c.suppressCloseNotify.Store(true)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UConn) Close() error {
|
func (c *UConn) Close() error {
|
||||||
if c.suppressCloseNotify.Load() {
|
|
||||||
return c.Conn.NetConn().Close()
|
|
||||||
}
|
|
||||||
timer := time.AfterFunc(tlsCloseTimeout, func() {
|
timer := time.AfterFunc(tlsCloseTimeout, func() {
|
||||||
c.Conn.NetConn().Close()
|
c.Conn.NetConn().Close()
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -2,9 +2,12 @@ package udp
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"reflect"
|
||||||
|
|
||||||
"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/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"
|
||||||
)
|
)
|
||||||
@@ -12,14 +15,40 @@ import (
|
|||||||
func init() {
|
func init() {
|
||||||
common.Must(internet.RegisterTransportDialer(protocolName,
|
common.Must(internet.RegisterTransportDialer(protocolName,
|
||||||
func(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (stat.Connection, error) {
|
func(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (stat.Connection, error) {
|
||||||
if streamSettings != nil && streamSettings.FinalMask != nil {
|
var sockopt *internet.SocketConfig
|
||||||
return streamSettings.FinalMask.DialUDP(ctx, dest)
|
if streamSettings != nil {
|
||||||
} else {
|
sockopt = streamSettings.SocketSettings
|
||||||
var sockopt *internet.SocketConfig
|
|
||||||
if streamSettings != nil && streamSettings.SocketSettings != nil {
|
|
||||||
sockopt = streamSettings.SocketSettings
|
|
||||||
}
|
|
||||||
return internet.DialSystem(ctx, dest, sockopt)
|
|
||||||
}
|
}
|
||||||
|
conn, err := internet.DialSystem(ctx, dest, sockopt)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if streamSettings != nil && 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,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return conn, nil
|
||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -58,15 +58,24 @@ func ListenUDP(ctx context.Context, address net.Address, port net.Port, streamSe
|
|||||||
}
|
}
|
||||||
|
|
||||||
var err error
|
var err error
|
||||||
if streamSettings.FinalMask != nil {
|
hub.conn, err = internet.ListenSystemPacket(ctx, &net.UDPAddr{
|
||||||
hub.conn, err = streamSettings.FinalMask.ListenPacket(ctx, &net.UDPAddr{IP: address.IP(), Port: int(port)})
|
IP: address.IP(),
|
||||||
} else {
|
Port: int(port),
|
||||||
hub.conn, err = internet.ListenSystemPacket(ctx, &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings)
|
}, sockopt)
|
||||||
}
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
raw := hub.conn
|
||||||
|
|
||||||
|
if streamSettings.UdpmaskManager != nil {
|
||||||
|
hub.conn, err = streamSettings.UdpmaskManager.WrapPacketConnServer(raw)
|
||||||
|
if err != nil {
|
||||||
|
raw.Close()
|
||||||
|
return nil, errors.New("mask err").Base(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
errors.LogInfo(ctx, "listening UDP on ", address, ":", port)
|
errors.LogInfo(ctx, "listening UDP on ", address, ":", port)
|
||||||
hub.udpConn, _ = hub.conn.(*net.UDPConn)
|
hub.udpConn, _ = hub.conn.(*net.UDPConn)
|
||||||
hub.cache = make(chan *udp.Packet, hub.capacity)
|
hub.cache = make(chan *udp.Packet, hub.capacity)
|
||||||
|
|||||||
@@ -48,16 +48,20 @@ func dialWebSocket(ctx context.Context, dest net.Destination, streamSettings *in
|
|||||||
|
|
||||||
dialer := &websocket.Dialer{
|
dialer := &websocket.Dialer{
|
||||||
NetDial: func(network, addr string) (net.Conn, error) {
|
NetDial: func(network, addr string) (net.Conn, error) {
|
||||||
var conn net.Conn
|
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||||
var err error
|
|
||||||
if streamSettings.FinalMask != nil {
|
|
||||||
conn, err = streamSettings.FinalMask.DialTCP(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").Base(err)
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if streamSettings.TcpmaskManager != nil {
|
||||||
|
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(conn)
|
||||||
|
if err != nil {
|
||||||
|
conn.Close()
|
||||||
|
return nil, errors.New("mask err").Base(err)
|
||||||
|
}
|
||||||
|
conn = newConn
|
||||||
|
}
|
||||||
|
|
||||||
return conn, err
|
return conn, err
|
||||||
},
|
},
|
||||||
ReadBufferSize: 4 * 1024,
|
ReadBufferSize: 4 * 1024,
|
||||||
@@ -75,15 +79,19 @@ func dialWebSocket(ctx context.Context, dest net.Destination, streamSettings *in
|
|||||||
if fingerprint := tls.GetFingerprint(tConfig.Fingerprint); fingerprint != nil {
|
if fingerprint := tls.GetFingerprint(tConfig.Fingerprint); fingerprint != nil {
|
||||||
dialer.NetDialTLSContext = func(_ context.Context, _, addr string) (net.Conn, error) {
|
dialer.NetDialTLSContext = func(_ context.Context, _, addr string) (net.Conn, error) {
|
||||||
// Like the NetDial in the dialer
|
// Like the NetDial in the dialer
|
||||||
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 {
|
||||||
return nil, errors.New("failed to dial to dest").Base(err)
|
errors.LogErrorInner(ctx, err, "failed to dial to "+addr)
|
||||||
|
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
|
||||||
}
|
}
|
||||||
|
|
||||||
// TLS and apply the handshake
|
// TLS and apply the handshake
|
||||||
|
|||||||
@@ -97,21 +97,29 @@ func ListenWS(ctx context.Context, address net.Address, port net.Port, streamSet
|
|||||||
}
|
}
|
||||||
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 WS) on ", address).Base(err)
|
||||||
|
}
|
||||||
|
errors.LogInfo(ctx, "listening unix domain socket(for WS) 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 WS) on ", address, ":", port).Base(err)
|
||||||
|
}
|
||||||
|
errors.LogInfo(ctx, "listening TCP(for WS) 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 WS) on ", address, ":", port).Base(err)
|
|
||||||
}
|
|
||||||
errors.LogInfo(ctx, "listening ", addr.Network(), "(for WS) 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")
|
||||||
|
|||||||
@@ -1,153 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
gotls "crypto/tls"
|
|
||||||
"net/http"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"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/reality"
|
|
||||||
"github.com/xtls/xray-core/transport/internet/tls"
|
|
||||||
"golang.org/x/net/http2"
|
|
||||||
)
|
|
||||||
|
|
||||||
type serviceTransport struct {
|
|
||||||
plain http.RoundTripper
|
|
||||||
secure http.RoundTripper
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *serviceTransport) RoundTrip(r *http.Request) (*http.Response, error) {
|
|
||||||
if r.URL.Scheme == "https" {
|
|
||||||
return t.secure.RoundTrip(r)
|
|
||||||
}
|
|
||||||
return t.plain.RoundTrip(r)
|
|
||||||
}
|
|
||||||
|
|
||||||
func newServiceClient(streamSettings *internet.MemoryStreamConfig, timeout time.Duration, maxConns int) *http.Client {
|
|
||||||
var (
|
|
||||||
tlsConfig *tls.Config
|
|
||||||
realityConfig *reality.Config
|
|
||||||
sockopt *internet.SocketConfig
|
|
||||||
fronting *net.Destination
|
|
||||||
)
|
|
||||||
if streamSettings != nil {
|
|
||||||
tlsConfig = tls.ConfigFromStreamSettings(streamSettings)
|
|
||||||
realityConfig = reality.ConfigFromStreamSettings(streamSettings)
|
|
||||||
sockopt = streamSettings.SocketSettings
|
|
||||||
fronting = streamSettings.Destination
|
|
||||||
}
|
|
||||||
overHTTP2 := allowsHTTP2(tlsConfig, realityConfig)
|
|
||||||
|
|
||||||
dial := func(ctx context.Context, addr string) (net.Conn, net.Destination, error) {
|
|
||||||
host, err := net.ParseDestination("tcp:" + addr)
|
|
||||||
if err != nil {
|
|
||||||
return nil, host, errors.New("bad address: ", addr).Base(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
target := host
|
|
||||||
if fronting != nil {
|
|
||||||
target.Address = fronting.Address
|
|
||||||
if fronting.Port != 0 {
|
|
||||||
target.Port = fronting.Port
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
var conn net.Conn
|
|
||||||
if streamSettings.FinalMask != nil {
|
|
||||||
conn, err = streamSettings.FinalMask.DialTCP(ctx, target)
|
|
||||||
} else {
|
|
||||||
conn, err = internet.DialSystem(ctx, target, sockopt)
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
return nil, host, errors.New("failed to dial to dest").Base(err)
|
|
||||||
}
|
|
||||||
return conn, host, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
dialPlain := func(ctx context.Context, network, addr string) (net.Conn, error) {
|
|
||||||
conn, _, err := dial(ctx, addr)
|
|
||||||
return conn, err
|
|
||||||
}
|
|
||||||
|
|
||||||
dialTLS := func(ctx context.Context, addr string) (net.Conn, error) {
|
|
||||||
conn, host, err := dial(ctx, addr)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if realityConfig != nil {
|
|
||||||
return reality.UClient(conn, realityConfig, ctx, host)
|
|
||||||
}
|
|
||||||
|
|
||||||
gotlsConfig := &gotls.Config{ServerName: host.Address.String()}
|
|
||||||
if tlsConfig != nil {
|
|
||||||
gotlsConfig = tlsConfig.GetTLSConfig(tls.WithDestination(host))
|
|
||||||
}
|
|
||||||
if len(gotlsConfig.NextProtos) != 1 {
|
|
||||||
if overHTTP2 {
|
|
||||||
gotlsConfig.NextProtos = []string{"h2"}
|
|
||||||
} else {
|
|
||||||
gotlsConfig.NextProtos = []string{"http/1.1"}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if tlsConfig != nil {
|
|
||||||
if fingerprint := tls.GetFingerprint(tlsConfig.Fingerprint); fingerprint != nil {
|
|
||||||
uconn := tls.UClient(conn, gotlsConfig, fingerprint)
|
|
||||||
if err := uconn.(*tls.UConn).HandshakeContext(ctx); err != nil {
|
|
||||||
conn.Close()
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return uconn, nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return tls.Client(conn, gotlsConfig), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var secure http.RoundTripper
|
|
||||||
if overHTTP2 {
|
|
||||||
secure = &http2.Transport{
|
|
||||||
DialTLSContext: func(ctx context.Context, network, addr string, cfg *gotls.Config) (net.Conn, error) {
|
|
||||||
return dialTLS(ctx, addr)
|
|
||||||
},
|
|
||||||
IdleConnTimeout: net.ConnIdleTimeout,
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
secure = &http.Transport{
|
|
||||||
DialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
|
||||||
return dialTLS(ctx, addr)
|
|
||||||
},
|
|
||||||
IdleConnTimeout: net.ConnIdleTimeout,
|
|
||||||
MaxIdleConns: maxConns,
|
|
||||||
MaxIdleConnsPerHost: maxConns,
|
|
||||||
MaxConnsPerHost: maxConns,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return &http.Client{
|
|
||||||
Transport: &serviceTransport{
|
|
||||||
plain: &http.Transport{
|
|
||||||
DialContext: dialPlain,
|
|
||||||
IdleConnTimeout: net.ConnIdleTimeout,
|
|
||||||
MaxIdleConns: maxConns,
|
|
||||||
MaxIdleConnsPerHost: maxConns,
|
|
||||||
MaxConnsPerHost: maxConns,
|
|
||||||
},
|
|
||||||
secure: secure,
|
|
||||||
},
|
|
||||||
Timeout: timeout,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func allowsHTTP2(tlsConfig *tls.Config, realityConfig *reality.Config) bool {
|
|
||||||
if realityConfig != nil {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
if tlsConfig == nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return len(tlsConfig.NextProtocol) == 1 && tlsConfig.NextProtocol[0] == "h2"
|
|
||||||
}
|
|
||||||
@@ -1,101 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
gotls "crypto/tls"
|
|
||||||
"net"
|
|
||||||
"net/http"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
xnet "github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
"github.com/xtls/xray-core/transport/internet/tls"
|
|
||||||
)
|
|
||||||
|
|
||||||
func recordingTLSListener(t *testing.T, sni *string, mu *sync.Mutex) net.Listener {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("listen: %v", err)
|
|
||||||
}
|
|
||||||
go func() {
|
|
||||||
for {
|
|
||||||
conn, err := ln.Accept()
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
cfg := &gotls.Config{
|
|
||||||
GetConfigForClient: func(hello *gotls.ClientHelloInfo) (*gotls.Config, error) {
|
|
||||||
mu.Lock()
|
|
||||||
*sni = hello.ServerName
|
|
||||||
mu.Unlock()
|
|
||||||
return nil, nil
|
|
||||||
},
|
|
||||||
}
|
|
||||||
tconn := gotls.Server(conn, cfg)
|
|
||||||
tconn.HandshakeContext(context.Background())
|
|
||||||
tconn.Close()
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
t.Cleanup(func() { ln.Close() })
|
|
||||||
return ln
|
|
||||||
}
|
|
||||||
|
|
||||||
func sniForSettings(t *testing.T, serverName string) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var (
|
|
||||||
sni string
|
|
||||||
mu sync.Mutex
|
|
||||||
)
|
|
||||||
ln := recordingTLSListener(t, &sni, &mu)
|
|
||||||
addr := ln.Addr().(*net.TCPAddr)
|
|
||||||
|
|
||||||
settings := &internet.MemoryStreamConfig{
|
|
||||||
ProtocolName: protocolName,
|
|
||||||
Destination: &xnet.Destination{
|
|
||||||
Address: xnet.ParseAddress(addr.IP.String()),
|
|
||||||
Port: xnet.Port(addr.Port),
|
|
||||||
Network: xnet.Network_TCP,
|
|
||||||
},
|
|
||||||
SecuritySettings: &tls.Config{ServerName: serverName},
|
|
||||||
}
|
|
||||||
|
|
||||||
prev := driveFilesURL
|
|
||||||
driveFilesURL = "https://www.googleapis.com/drive/v3/files"
|
|
||||||
defer func() { driveFilesURL = prev }()
|
|
||||||
|
|
||||||
client := newServiceClient(settings, 5*time.Second, 8)
|
|
||||||
req, err := http.NewRequest(http.MethodGet, driveFilesURL, nil)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("request: %v", err)
|
|
||||||
}
|
|
||||||
client.Do(req)
|
|
||||||
|
|
||||||
deadline := time.Now().Add(3 * time.Second)
|
|
||||||
for time.Now().Before(deadline) {
|
|
||||||
mu.Lock()
|
|
||||||
got := sni
|
|
||||||
mu.Unlock()
|
|
||||||
if got != "" {
|
|
||||||
return got
|
|
||||||
}
|
|
||||||
time.Sleep(5 * time.Millisecond)
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceSNIDefaultsToHost(t *testing.T) {
|
|
||||||
if got := sniForSettings(t, ""); got != "www.googleapis.com" {
|
|
||||||
t.Fatalf("SNI defaulted to %q, want the host www.googleapis.com, not address", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestServiceSNIOverride(t *testing.T) {
|
|
||||||
if got := sniForSettings(t, "www.google.com"); got != "www.google.com" {
|
|
||||||
t.Fatalf("explicit serverName gave SNI %q, want www.google.com", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,222 +0,0 @@
|
|||||||
// Code generated by protoc-gen-go. DO NOT EDIT.
|
|
||||||
// versions:
|
|
||||||
// protoc-gen-go v1.36.11
|
|
||||||
// protoc v6.33.5
|
|
||||||
// source: transport/internet/xdrive/config.proto
|
|
||||||
|
|
||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
|
|
||||||
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
|
|
||||||
reflect "reflect"
|
|
||||||
sync "sync"
|
|
||||||
unsafe "unsafe"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
// Verify that this generated code is sufficiently up-to-date.
|
|
||||||
_ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion)
|
|
||||||
// Verify that runtime/protoimpl is sufficiently up-to-date.
|
|
||||||
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
|
|
||||||
)
|
|
||||||
|
|
||||||
type Config struct {
|
|
||||||
state protoimpl.MessageState `protogen:"open.v1"`
|
|
||||||
RemoteFolder string `protobuf:"bytes,1,opt,name=remote_folder,json=remoteFolder,proto3" json:"remote_folder,omitempty"`
|
|
||||||
Service string `protobuf:"bytes,2,opt,name=service,proto3" json:"service,omitempty"`
|
|
||||||
Secrets []string `protobuf:"bytes,3,rep,name=secrets,proto3" json:"secrets,omitempty"`
|
|
||||||
SegmentBytes uint32 `protobuf:"varint,4,opt,name=segment_bytes,json=segmentBytes,proto3" json:"segment_bytes,omitempty"`
|
|
||||||
FlushIntervalMs uint32 `protobuf:"varint,5,opt,name=flush_interval_ms,json=flushIntervalMs,proto3" json:"flush_interval_ms,omitempty"`
|
|
||||||
PollIntervalMs uint32 `protobuf:"varint,6,opt,name=poll_interval_ms,json=pollIntervalMs,proto3" json:"poll_interval_ms,omitempty"`
|
|
||||||
MaxPollIntervalMs uint32 `protobuf:"varint,7,opt,name=max_poll_interval_ms,json=maxPollIntervalMs,proto3" json:"max_poll_interval_ms,omitempty"`
|
|
||||||
SessionTtlSeconds uint32 `protobuf:"varint,8,opt,name=session_ttl_seconds,json=sessionTtlSeconds,proto3" json:"session_ttl_seconds,omitempty"`
|
|
||||||
Concurrency uint32 `protobuf:"varint,9,opt,name=concurrency,proto3" json:"concurrency,omitempty"`
|
|
||||||
EagerWindowMs uint32 `protobuf:"varint,10,opt,name=eager_window_ms,json=eagerWindowMs,proto3" json:"eager_window_ms,omitempty"`
|
|
||||||
HoleTimeoutMs uint32 `protobuf:"varint,11,opt,name=hole_timeout_ms,json=holeTimeoutMs,proto3" json:"hole_timeout_ms,omitempty"`
|
|
||||||
Template string `protobuf:"bytes,12,opt,name=template,proto3" json:"template,omitempty"`
|
|
||||||
unknownFields protoimpl.UnknownFields
|
|
||||||
sizeCache protoimpl.SizeCache
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) Reset() {
|
|
||||||
*x = Config{}
|
|
||||||
mi := &file_transport_internet_xdrive_config_proto_msgTypes[0]
|
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
|
||||||
ms.StoreMessageInfo(mi)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) String() string {
|
|
||||||
return protoimpl.X.MessageStringOf(x)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (*Config) ProtoMessage() {}
|
|
||||||
|
|
||||||
func (x *Config) ProtoReflect() protoreflect.Message {
|
|
||||||
mi := &file_transport_internet_xdrive_config_proto_msgTypes[0]
|
|
||||||
if x != nil {
|
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
|
||||||
if ms.LoadMessageInfo() == nil {
|
|
||||||
ms.StoreMessageInfo(mi)
|
|
||||||
}
|
|
||||||
return ms
|
|
||||||
}
|
|
||||||
return mi.MessageOf(x)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
|
|
||||||
func (*Config) Descriptor() ([]byte, []int) {
|
|
||||||
return file_transport_internet_xdrive_config_proto_rawDescGZIP(), []int{0}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetRemoteFolder() string {
|
|
||||||
if x != nil {
|
|
||||||
return x.RemoteFolder
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetService() string {
|
|
||||||
if x != nil {
|
|
||||||
return x.Service
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetSecrets() []string {
|
|
||||||
if x != nil {
|
|
||||||
return x.Secrets
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetSegmentBytes() uint32 {
|
|
||||||
if x != nil {
|
|
||||||
return x.SegmentBytes
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetFlushIntervalMs() uint32 {
|
|
||||||
if x != nil {
|
|
||||||
return x.FlushIntervalMs
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetPollIntervalMs() uint32 {
|
|
||||||
if x != nil {
|
|
||||||
return x.PollIntervalMs
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetMaxPollIntervalMs() uint32 {
|
|
||||||
if x != nil {
|
|
||||||
return x.MaxPollIntervalMs
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetSessionTtlSeconds() uint32 {
|
|
||||||
if x != nil {
|
|
||||||
return x.SessionTtlSeconds
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetConcurrency() uint32 {
|
|
||||||
if x != nil {
|
|
||||||
return x.Concurrency
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetEagerWindowMs() uint32 {
|
|
||||||
if x != nil {
|
|
||||||
return x.EagerWindowMs
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetHoleTimeoutMs() uint32 {
|
|
||||||
if x != nil {
|
|
||||||
return x.HoleTimeoutMs
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *Config) GetTemplate() string {
|
|
||||||
if x != nil {
|
|
||||||
return x.Template
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
var File_transport_internet_xdrive_config_proto protoreflect.FileDescriptor
|
|
||||||
|
|
||||||
const file_transport_internet_xdrive_config_proto_rawDesc = "" +
|
|
||||||
"\n" +
|
|
||||||
"&transport/internet/xdrive/config.proto\x12\x1exray.transport.internet.xdrive\"\xcb\x03\n" +
|
|
||||||
"\x06Config\x12#\n" +
|
|
||||||
"\rremote_folder\x18\x01 \x01(\tR\fremoteFolder\x12\x18\n" +
|
|
||||||
"\aservice\x18\x02 \x01(\tR\aservice\x12\x18\n" +
|
|
||||||
"\asecrets\x18\x03 \x03(\tR\asecrets\x12#\n" +
|
|
||||||
"\rsegment_bytes\x18\x04 \x01(\rR\fsegmentBytes\x12*\n" +
|
|
||||||
"\x11flush_interval_ms\x18\x05 \x01(\rR\x0fflushIntervalMs\x12(\n" +
|
|
||||||
"\x10poll_interval_ms\x18\x06 \x01(\rR\x0epollIntervalMs\x12/\n" +
|
|
||||||
"\x14max_poll_interval_ms\x18\a \x01(\rR\x11maxPollIntervalMs\x12.\n" +
|
|
||||||
"\x13session_ttl_seconds\x18\b \x01(\rR\x11sessionTtlSeconds\x12 \n" +
|
|
||||||
"\vconcurrency\x18\t \x01(\rR\vconcurrency\x12&\n" +
|
|
||||||
"\x0feager_window_ms\x18\n" +
|
|
||||||
" \x01(\rR\reagerWindowMs\x12&\n" +
|
|
||||||
"\x0fhole_timeout_ms\x18\v \x01(\rR\rholeTimeoutMs\x12\x1a\n" +
|
|
||||||
"\btemplate\x18\f \x01(\tR\btemplateB5Z3github.com/xtls/xray-core/transport/internet/xdriveb\x06proto3"
|
|
||||||
|
|
||||||
var (
|
|
||||||
file_transport_internet_xdrive_config_proto_rawDescOnce sync.Once
|
|
||||||
file_transport_internet_xdrive_config_proto_rawDescData []byte
|
|
||||||
)
|
|
||||||
|
|
||||||
func file_transport_internet_xdrive_config_proto_rawDescGZIP() []byte {
|
|
||||||
file_transport_internet_xdrive_config_proto_rawDescOnce.Do(func() {
|
|
||||||
file_transport_internet_xdrive_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_transport_internet_xdrive_config_proto_rawDesc), len(file_transport_internet_xdrive_config_proto_rawDesc)))
|
|
||||||
})
|
|
||||||
return file_transport_internet_xdrive_config_proto_rawDescData
|
|
||||||
}
|
|
||||||
|
|
||||||
var file_transport_internet_xdrive_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
|
|
||||||
var file_transport_internet_xdrive_config_proto_goTypes = []any{
|
|
||||||
(*Config)(nil), // 0: xray.transport.internet.xdrive.Config
|
|
||||||
}
|
|
||||||
var file_transport_internet_xdrive_config_proto_depIdxs = []int32{
|
|
||||||
0, // [0:0] is the sub-list for method output_type
|
|
||||||
0, // [0:0] is the sub-list for method input_type
|
|
||||||
0, // [0:0] is the sub-list for extension type_name
|
|
||||||
0, // [0:0] is the sub-list for extension extendee
|
|
||||||
0, // [0:0] is the sub-list for field type_name
|
|
||||||
}
|
|
||||||
|
|
||||||
func init() { file_transport_internet_xdrive_config_proto_init() }
|
|
||||||
func file_transport_internet_xdrive_config_proto_init() {
|
|
||||||
if File_transport_internet_xdrive_config_proto != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
type x struct{}
|
|
||||||
out := protoimpl.TypeBuilder{
|
|
||||||
File: protoimpl.DescBuilder{
|
|
||||||
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
|
||||||
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_xdrive_config_proto_rawDesc), len(file_transport_internet_xdrive_config_proto_rawDesc)),
|
|
||||||
NumEnums: 0,
|
|
||||||
NumMessages: 1,
|
|
||||||
NumExtensions: 0,
|
|
||||||
NumServices: 0,
|
|
||||||
},
|
|
||||||
GoTypes: file_transport_internet_xdrive_config_proto_goTypes,
|
|
||||||
DependencyIndexes: file_transport_internet_xdrive_config_proto_depIdxs,
|
|
||||||
MessageInfos: file_transport_internet_xdrive_config_proto_msgTypes,
|
|
||||||
}.Build()
|
|
||||||
File_transport_internet_xdrive_config_proto = out.File
|
|
||||||
file_transport_internet_xdrive_config_proto_goTypes = nil
|
|
||||||
file_transport_internet_xdrive_config_proto_depIdxs = nil
|
|
||||||
}
|
|
||||||
@@ -1,19 +0,0 @@
|
|||||||
syntax = "proto3";
|
|
||||||
|
|
||||||
package xray.transport.internet.xdrive;
|
|
||||||
option go_package = "github.com/xtls/xray-core/transport/internet/xdrive";
|
|
||||||
|
|
||||||
message Config {
|
|
||||||
string remote_folder = 1;
|
|
||||||
string service = 2;
|
|
||||||
repeated string secrets = 3;
|
|
||||||
uint32 segment_bytes = 4;
|
|
||||||
uint32 flush_interval_ms = 5;
|
|
||||||
uint32 poll_interval_ms = 6;
|
|
||||||
uint32 max_poll_interval_ms = 7;
|
|
||||||
uint32 session_ttl_seconds = 8;
|
|
||||||
uint32 concurrency = 9;
|
|
||||||
uint32 eager_window_ms = 10;
|
|
||||||
uint32 hole_timeout_ms = 11;
|
|
||||||
string template = 12;
|
|
||||||
}
|
|
||||||
@@ -1,139 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"os"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
)
|
|
||||||
|
|
||||||
var placeholderAddr = &net.TCPAddr{IP: net.IP{127, 0, 0, 1}, Port: 0}
|
|
||||||
|
|
||||||
type Conn struct {
|
|
||||||
cancel context.CancelFunc
|
|
||||||
writer *walWriter
|
|
||||||
reader *walReader
|
|
||||||
onClose func()
|
|
||||||
|
|
||||||
readBuf []byte
|
|
||||||
|
|
||||||
deadlineMu sync.Mutex
|
|
||||||
readDeadline time.Time
|
|
||||||
writeDeadline time.Time
|
|
||||||
|
|
||||||
closeOnce sync.Once
|
|
||||||
closeErr error
|
|
||||||
}
|
|
||||||
|
|
||||||
func newConn(ctx context.Context, storage Storage, writePrefix, readPrefix string, p params, onClose func()) *Conn {
|
|
||||||
ctx, cancel := context.WithCancel(ctx)
|
|
||||||
return &Conn{
|
|
||||||
cancel: cancel,
|
|
||||||
writer: newWALWriter(ctx, storage, writePrefix, p),
|
|
||||||
reader: newWALReader(ctx, storage, readPrefix, p),
|
|
||||||
onClose: onClose,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) Read(b []byte) (int, error) {
|
|
||||||
if len(c.readBuf) == 0 {
|
|
||||||
data, err := c.receive()
|
|
||||||
if err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
c.readBuf = data
|
|
||||||
}
|
|
||||||
n := copy(b, c.readBuf)
|
|
||||||
c.readBuf = c.readBuf[n:]
|
|
||||||
return n, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) receive() ([]byte, error) {
|
|
||||||
deadline := c.getDeadline(true)
|
|
||||||
if deadline.IsZero() {
|
|
||||||
data, ok := <-c.reader.ch
|
|
||||||
if !ok {
|
|
||||||
return nil, c.reader.Err()
|
|
||||||
}
|
|
||||||
return data, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if !time.Now().Before(deadline) {
|
|
||||||
return nil, os.ErrDeadlineExceeded
|
|
||||||
}
|
|
||||||
timer := time.NewTimer(time.Until(deadline))
|
|
||||||
defer timer.Stop()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case data, ok := <-c.reader.ch:
|
|
||||||
if !ok {
|
|
||||||
return nil, c.reader.Err()
|
|
||||||
}
|
|
||||||
return data, nil
|
|
||||||
case <-timer.C:
|
|
||||||
return nil, os.ErrDeadlineExceeded
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) Write(b []byte) (int, error) {
|
|
||||||
if deadline := c.getDeadline(false); !deadline.IsZero() && !time.Now().Before(deadline) {
|
|
||||||
return 0, os.ErrDeadlineExceeded
|
|
||||||
}
|
|
||||||
n, err := c.writer.Write(b)
|
|
||||||
if err == nil {
|
|
||||||
c.reader.Wake()
|
|
||||||
}
|
|
||||||
return n, err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) Close() error {
|
|
||||||
c.closeOnce.Do(func() {
|
|
||||||
c.closeErr = c.writer.Close()
|
|
||||||
c.cancel()
|
|
||||||
if c.onClose != nil {
|
|
||||||
c.onClose()
|
|
||||||
}
|
|
||||||
})
|
|
||||||
return c.closeErr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) LocalAddr() net.Addr {
|
|
||||||
return placeholderAddr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) RemoteAddr() net.Addr {
|
|
||||||
return placeholderAddr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) getDeadline(read bool) time.Time {
|
|
||||||
c.deadlineMu.Lock()
|
|
||||||
defer c.deadlineMu.Unlock()
|
|
||||||
if read {
|
|
||||||
return c.readDeadline
|
|
||||||
}
|
|
||||||
return c.writeDeadline
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) SetDeadline(t time.Time) error {
|
|
||||||
c.deadlineMu.Lock()
|
|
||||||
defer c.deadlineMu.Unlock()
|
|
||||||
c.readDeadline = t
|
|
||||||
c.writeDeadline = t
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) SetReadDeadline(t time.Time) error {
|
|
||||||
c.deadlineMu.Lock()
|
|
||||||
defer c.deadlineMu.Unlock()
|
|
||||||
c.readDeadline = t
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) SetWriteDeadline(t time.Time) error {
|
|
||||||
c.deadlineMu.Lock()
|
|
||||||
defer c.deadlineMu.Unlock()
|
|
||||||
c.writeDeadline = t
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,575 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"encoding/base64"
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net/http"
|
|
||||||
"net/url"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/dice"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
flatSeparator = "~"
|
|
||||||
|
|
||||||
driveBoundary = "xdrive-boundary"
|
|
||||||
drivePageSize = 1000
|
|
||||||
driveMaxAttempts = 8
|
|
||||||
driveMaxInflight = 32
|
|
||||||
driveInlineLimit = 12000
|
|
||||||
driveTimeout = 60 * time.Second
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
driveMaxBackoff = 8 * time.Second
|
|
||||||
driveTokenURL = "https://oauth2.googleapis.com/token"
|
|
||||||
driveFilesURL = "https://www.googleapis.com/drive/v3/files"
|
|
||||||
driveUploadURL = "https://www.googleapis.com/upload/drive/v3/files?uploadType=multipart&fields=id,name"
|
|
||||||
driveInitialBackoff = 200 * time.Millisecond
|
|
||||||
)
|
|
||||||
|
|
||||||
type driveStorage struct {
|
|
||||||
folder string
|
|
||||||
clientID string
|
|
||||||
clientSecret string
|
|
||||||
refreshToken string
|
|
||||||
client *http.Client
|
|
||||||
tokenURL string
|
|
||||||
filesURL string
|
|
||||||
uploadURL string
|
|
||||||
backoff time.Duration
|
|
||||||
|
|
||||||
tokenMu sync.Mutex
|
|
||||||
token string
|
|
||||||
tokenExpiry time.Time
|
|
||||||
|
|
||||||
inflight chan struct{}
|
|
||||||
|
|
||||||
idMu sync.Mutex
|
|
||||||
ids map[string]string
|
|
||||||
}
|
|
||||||
|
|
||||||
func newDriveStorage(streamSettings *internet.MemoryStreamConfig, config *Config) (*driveStorage, error) {
|
|
||||||
if config.RemoteFolder == "" {
|
|
||||||
return nil, errors.New(`empty "remoteFolder", it must be a Google Drive folder id`)
|
|
||||||
}
|
|
||||||
if len(config.Secrets) != 3 {
|
|
||||||
return nil, errors.New("Google Drive needs 3 secrets in order of ClientID, ClientSecret, RefreshToken")
|
|
||||||
}
|
|
||||||
for i, secret := range config.Secrets {
|
|
||||||
if secret == "" {
|
|
||||||
return nil, errors.New("Google Drive secret ", i, " is empty")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return &driveStorage{
|
|
||||||
folder: config.RemoteFolder,
|
|
||||||
clientID: config.Secrets[0],
|
|
||||||
clientSecret: config.Secrets[1],
|
|
||||||
refreshToken: config.Secrets[2],
|
|
||||||
client: newServiceClient(streamSettings, driveTimeout, driveMaxInflight),
|
|
||||||
tokenURL: driveTokenURL,
|
|
||||||
filesURL: driveFilesURL,
|
|
||||||
uploadURL: driveUploadURL,
|
|
||||||
backoff: driveInitialBackoff,
|
|
||||||
inflight: make(chan struct{}, driveMaxInflight),
|
|
||||||
ids: make(map[string]string),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func flatten(name string) string {
|
|
||||||
return strings.ReplaceAll(name, "/", flatSeparator)
|
|
||||||
}
|
|
||||||
|
|
||||||
func quoteDriveValue(value string) string {
|
|
||||||
return strings.NewReplacer(`\`, `\\`, `'`, `\'`).Replace(value)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) accessToken(ctx context.Context) (string, error) {
|
|
||||||
s.tokenMu.Lock()
|
|
||||||
defer s.tokenMu.Unlock()
|
|
||||||
|
|
||||||
if s.token != "" && time.Now().Before(s.tokenExpiry) {
|
|
||||||
return s.token, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
form := url.Values{
|
|
||||||
"client_id": {s.clientID},
|
|
||||||
"client_secret": {s.clientSecret},
|
|
||||||
"refresh_token": {s.refreshToken},
|
|
||||||
"grant_type": {"refresh_token"},
|
|
||||||
}
|
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, s.tokenURL, strings.NewReader(form.Encode()))
|
|
||||||
if err != nil {
|
|
||||||
return "", errors.New("failed to build the token request").Base(err)
|
|
||||||
}
|
|
||||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
|
||||||
|
|
||||||
resp, err := s.client.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
return "", errors.New("failed to refresh the access token").Base(err)
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
|
|
||||||
if err != nil {
|
|
||||||
return "", errors.New("failed to read the token response").Base(err)
|
|
||||||
}
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
|
||||||
return "", errors.New("the token endpoint answered ", resp.StatusCode, ": ", string(body))
|
|
||||||
}
|
|
||||||
|
|
||||||
var parsed struct {
|
|
||||||
AccessToken string `json:"access_token"`
|
|
||||||
ExpiresIn int64 `json:"expires_in"`
|
|
||||||
}
|
|
||||||
if err := json.Unmarshal(body, &parsed); err != nil {
|
|
||||||
return "", errors.New("failed to parse the token response").Base(err)
|
|
||||||
}
|
|
||||||
if parsed.AccessToken == "" {
|
|
||||||
return "", errors.New("the token endpoint returned no access token")
|
|
||||||
}
|
|
||||||
|
|
||||||
lifetime := parsed.ExpiresIn
|
|
||||||
if lifetime > 60 {
|
|
||||||
lifetime -= 60
|
|
||||||
}
|
|
||||||
s.token = parsed.AccessToken
|
|
||||||
s.tokenExpiry = time.Now().Add(time.Duration(lifetime) * time.Second)
|
|
||||||
return s.token, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func jitter(backoff time.Duration) time.Duration {
|
|
||||||
half := backoff / 2
|
|
||||||
if half <= 0 {
|
|
||||||
return backoff
|
|
||||||
}
|
|
||||||
return half + time.Duration(dice.Roll(int(half)))
|
|
||||||
}
|
|
||||||
|
|
||||||
func rateLimited(payload []byte) bool {
|
|
||||||
var parsed struct {
|
|
||||||
Error struct {
|
|
||||||
Status string `json:"status"`
|
|
||||||
Errors []struct {
|
|
||||||
Reason string `json:"reason"`
|
|
||||||
} `json:"errors"`
|
|
||||||
} `json:"error"`
|
|
||||||
}
|
|
||||||
if json.Unmarshal(payload, &parsed) != nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
for _, item := range parsed.Error.Errors {
|
|
||||||
switch item.Reason {
|
|
||||||
case "rateLimitExceeded", "userRateLimitExceeded", "sharingRateLimitExceeded":
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return parsed.Error.Status == "RESOURCE_EXHAUSTED"
|
|
||||||
}
|
|
||||||
|
|
||||||
func retryableStatus(status int) bool {
|
|
||||||
switch status {
|
|
||||||
case http.StatusTooManyRequests, http.StatusInternalServerError,
|
|
||||||
http.StatusBadGateway, http.StatusServiceUnavailable, http.StatusGatewayTimeout:
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) do(ctx context.Context, method, target, contentType string, body []byte) (int, []byte, error) {
|
|
||||||
backoff := s.backoff
|
|
||||||
var lastErr error
|
|
||||||
|
|
||||||
for attempt := 0; attempt < driveMaxAttempts; attempt++ {
|
|
||||||
if attempt > 0 {
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
return 0, nil, ctx.Err()
|
|
||||||
case <-time.After(jitter(backoff)):
|
|
||||||
}
|
|
||||||
backoff *= 2
|
|
||||||
if backoff > driveMaxBackoff {
|
|
||||||
backoff = driveMaxBackoff
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
token, err := s.accessToken(ctx)
|
|
||||||
if err != nil {
|
|
||||||
lastErr = err
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
select {
|
|
||||||
case s.inflight <- struct{}{}:
|
|
||||||
case <-ctx.Done():
|
|
||||||
return 0, nil, ctx.Err()
|
|
||||||
}
|
|
||||||
|
|
||||||
var reader io.Reader
|
|
||||||
if body != nil {
|
|
||||||
reader = bytes.NewReader(body)
|
|
||||||
}
|
|
||||||
req, err := http.NewRequestWithContext(ctx, method, target, reader)
|
|
||||||
if err != nil {
|
|
||||||
return 0, nil, errors.New("failed to build a Drive request").Base(err)
|
|
||||||
}
|
|
||||||
req.Header.Set("Authorization", "Bearer "+token)
|
|
||||||
if contentType != "" {
|
|
||||||
req.Header.Set("Content-Type", contentType)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp, err := s.client.Do(req)
|
|
||||||
<-s.inflight
|
|
||||||
if err != nil {
|
|
||||||
if ctx.Err() != nil {
|
|
||||||
return 0, nil, ctx.Err()
|
|
||||||
}
|
|
||||||
lastErr = errors.New("Drive request failed").Base(err)
|
|
||||||
errors.LogWarningInner(ctx, err, "retrying a failed Drive request, attempt ",
|
|
||||||
attempt+1, " of ", driveMaxAttempts)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
payload, err := io.ReadAll(resp.Body)
|
|
||||||
resp.Body.Close()
|
|
||||||
if err != nil {
|
|
||||||
if ctx.Err() != nil {
|
|
||||||
return 0, nil, ctx.Err()
|
|
||||||
}
|
|
||||||
lastErr = errors.New("failed to read the Drive response").Base(err)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if resp.StatusCode == http.StatusUnauthorized {
|
|
||||||
s.invalidateToken()
|
|
||||||
lastErr = errors.New("Drive rejected the access token")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if resp.StatusCode == http.StatusForbidden && rateLimited(payload) {
|
|
||||||
lastErr = errors.New("Drive is rate limiting: ", string(payload))
|
|
||||||
errors.LogWarning(ctx, "rate limited by Drive, attempt ",
|
|
||||||
attempt+1, " of ", driveMaxAttempts)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if retryableStatus(resp.StatusCode) {
|
|
||||||
lastErr = errors.New("Drive answered ", resp.StatusCode, ": ", string(payload))
|
|
||||||
errors.LogWarning(ctx, "retrying after Drive answered ", resp.StatusCode,
|
|
||||||
", attempt ", attempt+1, " of ", driveMaxAttempts)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
return resp.StatusCode, payload, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return 0, nil, lastErr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) invalidateToken() {
|
|
||||||
s.tokenMu.Lock()
|
|
||||||
s.token = ""
|
|
||||||
s.tokenMu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) rememberID(name, id string) {
|
|
||||||
s.idMu.Lock()
|
|
||||||
s.ids[name] = id
|
|
||||||
s.idMu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) forgetID(name string) {
|
|
||||||
s.idMu.Lock()
|
|
||||||
delete(s.ids, name)
|
|
||||||
s.idMu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) cachedID(name string) (string, bool) {
|
|
||||||
s.idMu.Lock()
|
|
||||||
defer s.idMu.Unlock()
|
|
||||||
id, ok := s.ids[name]
|
|
||||||
return id, ok
|
|
||||||
}
|
|
||||||
|
|
||||||
type driveFile struct {
|
|
||||||
id string
|
|
||||||
description string
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) query(ctx context.Context, condition string) (map[string]driveFile, error) {
|
|
||||||
found := make(map[string]driveFile)
|
|
||||||
pageToken := ""
|
|
||||||
|
|
||||||
for {
|
|
||||||
params := url.Values{
|
|
||||||
"q": {"'" + quoteDriveValue(s.folder) + "' in parents and trashed = false and " + condition},
|
|
||||||
"fields": {"nextPageToken,files(id,name,description)"},
|
|
||||||
"pageSize": {fmt.Sprint(drivePageSize)},
|
|
||||||
"supportsAllDrives": {"true"},
|
|
||||||
"includeItemsFromAllDrives": {"true"},
|
|
||||||
}
|
|
||||||
if pageToken != "" {
|
|
||||||
params.Set("pageToken", pageToken)
|
|
||||||
}
|
|
||||||
|
|
||||||
status, payload, err := s.do(ctx, http.MethodGet, s.filesURL+"?"+params.Encode(), "", nil)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if status != http.StatusOK {
|
|
||||||
return nil, errors.New("Drive listing answered ", status, ": ", string(payload))
|
|
||||||
}
|
|
||||||
|
|
||||||
var parsed struct {
|
|
||||||
NextPageToken string `json:"nextPageToken"`
|
|
||||||
Files []struct {
|
|
||||||
ID string `json:"id"`
|
|
||||||
Name string `json:"name"`
|
|
||||||
Description string `json:"description"`
|
|
||||||
} `json:"files"`
|
|
||||||
}
|
|
||||||
if err := json.Unmarshal(payload, &parsed); err != nil {
|
|
||||||
return nil, errors.New("failed to parse the Drive listing").Base(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, file := range parsed.Files {
|
|
||||||
found[file.Name] = driveFile{id: file.ID, description: file.Description}
|
|
||||||
s.rememberID(file.Name, file.ID)
|
|
||||||
}
|
|
||||||
|
|
||||||
pageToken = parsed.NextPageToken
|
|
||||||
if pageToken == "" {
|
|
||||||
return found, nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) resolveID(ctx context.Context, flat string) (string, error) {
|
|
||||||
if id, ok := s.cachedID(flat); ok {
|
|
||||||
return id, nil
|
|
||||||
}
|
|
||||||
found, err := s.query(ctx, "name = '"+quoteDriveValue(flat)+"'")
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
if file, ok := found[flat]; ok {
|
|
||||||
return file.id, nil
|
|
||||||
}
|
|
||||||
return "", errNotFound
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) Put(ctx context.Context, name string, data []byte) error {
|
|
||||||
if len(data) <= driveInlineLimit {
|
|
||||||
return s.putInline(ctx, name, data)
|
|
||||||
}
|
|
||||||
return s.putMedia(ctx, name, data)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) putInline(ctx context.Context, name string, data []byte) error {
|
|
||||||
flat := flatten(name)
|
|
||||||
|
|
||||||
body, err := json.Marshal(map[string]interface{}{
|
|
||||||
"name": flat,
|
|
||||||
"parents": []string{s.folder},
|
|
||||||
"description": base64.StdEncoding.EncodeToString(data),
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return errors.New("failed to build the inline metadata").Base(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
status, payload, err := s.do(ctx, http.MethodPost, s.filesURL+"?fields=id",
|
|
||||||
"application/json; charset=UTF-8", body)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if status != http.StatusOK {
|
|
||||||
return errors.New("Drive rejected the inline upload of ", name,
|
|
||||||
" with ", status, ": ", string(payload))
|
|
||||||
}
|
|
||||||
|
|
||||||
var parsed struct {
|
|
||||||
ID string `json:"id"`
|
|
||||||
}
|
|
||||||
if err := json.Unmarshal(payload, &parsed); err != nil {
|
|
||||||
return errors.New("failed to parse the Drive upload response").Base(err)
|
|
||||||
}
|
|
||||||
if parsed.ID != "" {
|
|
||||||
s.rememberID(flat, parsed.ID)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) putMedia(ctx context.Context, name string, data []byte) error {
|
|
||||||
flat := flatten(name)
|
|
||||||
|
|
||||||
metadata, err := json.Marshal(map[string]interface{}{
|
|
||||||
"name": flat,
|
|
||||||
"parents": []string{s.folder},
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return errors.New("failed to build the upload metadata").Base(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var body bytes.Buffer
|
|
||||||
fmt.Fprintf(&body, "--%s\r\nContent-Type: application/json; charset=UTF-8\r\n\r\n", driveBoundary)
|
|
||||||
body.Write(metadata)
|
|
||||||
fmt.Fprintf(&body, "\r\n--%s\r\nContent-Type: application/octet-stream\r\n\r\n", driveBoundary)
|
|
||||||
body.Write(data)
|
|
||||||
fmt.Fprintf(&body, "\r\n--%s--\r\n", driveBoundary)
|
|
||||||
|
|
||||||
status, payload, err := s.do(ctx, http.MethodPost, s.uploadURL,
|
|
||||||
"multipart/related; boundary="+driveBoundary, body.Bytes())
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if status != http.StatusOK {
|
|
||||||
return errors.New("Drive upload of ", name, " answered ", status, ": ", string(payload))
|
|
||||||
}
|
|
||||||
|
|
||||||
var parsed struct {
|
|
||||||
ID string `json:"id"`
|
|
||||||
}
|
|
||||||
if err := json.Unmarshal(payload, &parsed); err != nil {
|
|
||||||
return errors.New("failed to parse the Drive upload response").Base(err)
|
|
||||||
}
|
|
||||||
if parsed.ID != "" {
|
|
||||||
s.rememberID(flat, parsed.ID)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) Get(ctx context.Context, name string) ([]byte, error) {
|
|
||||||
flat := flatten(name)
|
|
||||||
id, err := s.resolveID(ctx, flat)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
status, payload, err := s.do(ctx, http.MethodGet,
|
|
||||||
s.filesURL+"/"+url.PathEscape(id)+"?alt=media&supportsAllDrives=true", "", nil)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
switch status {
|
|
||||||
case http.StatusOK:
|
|
||||||
if len(payload) > 0 {
|
|
||||||
return payload, nil
|
|
||||||
}
|
|
||||||
return s.getInline(ctx, id)
|
|
||||||
case http.StatusNotFound:
|
|
||||||
s.forgetID(flat)
|
|
||||||
return nil, errNotFound
|
|
||||||
default:
|
|
||||||
return nil, errors.New("Drive download of ", name, " answered ", status, ": ", string(payload))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) getInline(ctx context.Context, id string) ([]byte, error) {
|
|
||||||
status, payload, err := s.do(ctx, http.MethodGet,
|
|
||||||
s.filesURL+"/"+url.PathEscape(id)+"?fields=description&supportsAllDrives=true", "", nil)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if status != http.StatusOK {
|
|
||||||
return nil, errors.New("Drive answered ", status, " for inline data: ", string(payload))
|
|
||||||
}
|
|
||||||
|
|
||||||
var parsed struct {
|
|
||||||
Description string `json:"description"`
|
|
||||||
}
|
|
||||||
if err := json.Unmarshal(payload, &parsed); err != nil {
|
|
||||||
return nil, errors.New("failed to parse the inline data").Base(err)
|
|
||||||
}
|
|
||||||
if parsed.Description == "" {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
data, err := base64.StdEncoding.DecodeString(parsed.Description)
|
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("the inline data is not valid base64").Base(err)
|
|
||||||
}
|
|
||||||
return data, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) deleteID(ctx context.Context, flat, id string) error {
|
|
||||||
status, payload, err := s.do(ctx, http.MethodDelete,
|
|
||||||
s.filesURL+"/"+url.PathEscape(id)+"?supportsAllDrives=true", "", nil)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
s.forgetID(flat)
|
|
||||||
switch status {
|
|
||||||
case http.StatusOK, http.StatusNoContent, http.StatusNotFound:
|
|
||||||
return nil
|
|
||||||
default:
|
|
||||||
return errors.New("Drive deletion of ", flat, " answered ", status, ": ", string(payload))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) Delete(ctx context.Context, name string) error {
|
|
||||||
flat := flatten(name)
|
|
||||||
|
|
||||||
if id, err := s.resolveID(ctx, flat); err == nil {
|
|
||||||
if err := s.deleteID(ctx, flat, id); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
} else if err != errNotFound {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
children, err := s.query(ctx, "name contains '"+quoteDriveValue(flat+flatSeparator)+"'")
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
for childName, file := range children {
|
|
||||||
if err := s.deleteID(ctx, childName, file.id); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) List(ctx context.Context, prefix string) ([]Entry, error) {
|
|
||||||
flat := flatten(prefix) + flatSeparator
|
|
||||||
|
|
||||||
found, err := s.query(ctx, "name contains '"+quoteDriveValue(flat)+"'")
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
seen := make(map[string]bool, len(found))
|
|
||||||
entries := make([]Entry, 0, len(found))
|
|
||||||
for name, file := range found {
|
|
||||||
rest := strings.TrimPrefix(name, flat)
|
|
||||||
if rest == "" {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
direct := true
|
|
||||||
if cut := strings.Index(rest, flatSeparator); cut >= 0 {
|
|
||||||
rest = rest[:cut]
|
|
||||||
direct = false
|
|
||||||
}
|
|
||||||
if seen[rest] {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
seen[rest] = true
|
|
||||||
|
|
||||||
entry := Entry{Name: rest}
|
|
||||||
if direct && file.description != "" {
|
|
||||||
if data, err := base64.StdEncoding.DecodeString(file.description); err == nil {
|
|
||||||
entry.Inline = data
|
|
||||||
}
|
|
||||||
}
|
|
||||||
entries = append(entries, entry)
|
|
||||||
}
|
|
||||||
return entries, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *driveStorage) Close() error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,345 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"crypto/rand"
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
"strconv"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
)
|
|
||||||
|
|
||||||
const liveSecretsEnv = "XRAY_XDRIVE_DRIVE_SECRETS"
|
|
||||||
|
|
||||||
func liveDriveConfig(t *testing.T) *Config {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
path := os.Getenv(liveSecretsEnv)
|
|
||||||
if path == "" {
|
|
||||||
t.Skipf("set %s to run this test", liveSecretsEnv)
|
|
||||||
}
|
|
||||||
|
|
||||||
payload, err := os.ReadFile(path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("reading %s: %v", path, err)
|
|
||||||
}
|
|
||||||
var secrets struct {
|
|
||||||
Folder string `json:"folder"`
|
|
||||||
ClientID string `json:"client_id"`
|
|
||||||
ClientSecret string `json:"client_secret"`
|
|
||||||
RefreshToken string `json:"refresh_token"`
|
|
||||||
}
|
|
||||||
if err := json.Unmarshal(payload, &secrets); err != nil {
|
|
||||||
t.Fatalf("parsing %s: %v", path, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
config := &Config{
|
|
||||||
RemoteFolder: secrets.Folder,
|
|
||||||
Service: "Google Drive",
|
|
||||||
Secrets: []string{secrets.ClientID, secrets.ClientSecret, secrets.RefreshToken},
|
|
||||||
SegmentBytes: 256 * 1024,
|
|
||||||
FlushIntervalMs: 100,
|
|
||||||
PollIntervalMs: 500,
|
|
||||||
MaxPollIntervalMs: 2000,
|
|
||||||
SessionTtlSeconds: 120,
|
|
||||||
}
|
|
||||||
if raw := os.Getenv("XRAY_XDRIVE_LIVE_SEGMENT"); raw != "" {
|
|
||||||
config.SegmentBytes = uint32(envInt(t, "XRAY_XDRIVE_LIVE_SEGMENT"))
|
|
||||||
}
|
|
||||||
if raw := os.Getenv("XRAY_XDRIVE_LIVE_CONCURRENCY"); raw != "" {
|
|
||||||
config.Concurrency = uint32(envInt(t, "XRAY_XDRIVE_LIVE_CONCURRENCY"))
|
|
||||||
}
|
|
||||||
|
|
||||||
storage, err := newDriveStorage(nil, config)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newDriveStorage: %v", err)
|
|
||||||
}
|
|
||||||
defer storage.Close()
|
|
||||||
for _, dir := range []string{sessionsDir, streamsDir} {
|
|
||||||
if err := storage.Delete(context.Background(), dir); err != nil {
|
|
||||||
t.Fatalf("clearing %s: %v", dir, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return config
|
|
||||||
}
|
|
||||||
|
|
||||||
func envInt(t *testing.T, name string) int {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
parsed, err := strconv.Atoi(os.Getenv(name))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("%s: %v", name, err)
|
|
||||||
}
|
|
||||||
return parsed
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLiveDriveStorage(t *testing.T) {
|
|
||||||
config := liveDriveConfig(t)
|
|
||||||
storage, err := newDriveStorage(nil, config)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newDriveStorage: %v", err)
|
|
||||||
}
|
|
||||||
defer storage.Close()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
name := "streams/livetest/c2s/000000000.seg"
|
|
||||||
payload := []byte("xdrive over a real remote storage service")
|
|
||||||
defer storage.Delete(ctx, "streams/livetest")
|
|
||||||
|
|
||||||
start := time.Now()
|
|
||||||
if err := storage.Put(ctx, name, payload); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
t.Logf("Put took %v", time.Since(start))
|
|
||||||
|
|
||||||
start = time.Now()
|
|
||||||
names, err := storage.List(ctx, "streams/livetest/c2s")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List: %v", err)
|
|
||||||
}
|
|
||||||
t.Logf("List took %v", time.Since(start))
|
|
||||||
if len(names) != 1 || names[0].Name != "000000000.seg" {
|
|
||||||
t.Fatalf("List returned %v, want one segment", names)
|
|
||||||
}
|
|
||||||
|
|
||||||
start = time.Now()
|
|
||||||
got, err := storage.Get(ctx, name)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Get: %v", err)
|
|
||||||
}
|
|
||||||
t.Logf("Get took %v", time.Since(start))
|
|
||||||
if !bytes.Equal(got, payload) {
|
|
||||||
t.Fatalf("Get returned %q, want %q", got, payload)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := storage.Get(ctx, "streams/livetest/c2s/000000009.seg"); err != errNotFound {
|
|
||||||
t.Fatalf("Get returned %v, want errNotFound", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := storage.Delete(ctx, "streams/livetest"); err != nil {
|
|
||||||
t.Fatalf("Delete: %v", err)
|
|
||||||
}
|
|
||||||
names, err = storage.List(ctx, "streams/livetest/c2s")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List after delete: %v", err)
|
|
||||||
}
|
|
||||||
if len(names) != 0 {
|
|
||||||
t.Fatalf("List after delete returned %v", names)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLiveDriveTransport(t *testing.T) {
|
|
||||||
config := liveDriveConfig(t)
|
|
||||||
streamSettings := &internet.MemoryStreamConfig{
|
|
||||||
ProtocolName: protocolName,
|
|
||||||
ProtocolSettings: config,
|
|
||||||
}
|
|
||||||
|
|
||||||
client, server, cleanup := pairWith(t, streamSettings)
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
start := time.Now()
|
|
||||||
if _, err := client.Write([]byte("ping")); err != nil {
|
|
||||||
t.Fatalf("client write: %v", err)
|
|
||||||
}
|
|
||||||
expectRead(t, server, "ping")
|
|
||||||
t.Logf("client to server round took %v", time.Since(start))
|
|
||||||
|
|
||||||
start = time.Now()
|
|
||||||
if _, err := server.Write([]byte("pong")); err != nil {
|
|
||||||
t.Fatalf("server write: %v", err)
|
|
||||||
}
|
|
||||||
expectRead(t, client, "pong")
|
|
||||||
t.Logf("server to client round took %v", time.Since(start))
|
|
||||||
|
|
||||||
size := 400000
|
|
||||||
if raw := os.Getenv("XRAY_XDRIVE_LIVE_BYTES"); raw != "" {
|
|
||||||
parsed, err := strconv.Atoi(raw)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("XRAY_XDRIVE_LIVE_BYTES: %v", err)
|
|
||||||
}
|
|
||||||
size = parsed
|
|
||||||
}
|
|
||||||
payload := make([]byte, size)
|
|
||||||
if _, err := rand.Read(payload); err != nil {
|
|
||||||
t.Fatalf("rand: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
start = time.Now()
|
|
||||||
go func() {
|
|
||||||
client.Write(payload)
|
|
||||||
}()
|
|
||||||
if err := server.SetReadDeadline(time.Now().Add(5 * time.Minute)); err != nil {
|
|
||||||
t.Fatalf("SetReadDeadline: %v", err)
|
|
||||||
}
|
|
||||||
got := make([]byte, len(payload))
|
|
||||||
if _, err := io.ReadFull(server, got); err != nil {
|
|
||||||
t.Fatalf("ReadFull: %v", err)
|
|
||||||
}
|
|
||||||
elapsed := time.Since(start)
|
|
||||||
if !bytes.Equal(got, payload) {
|
|
||||||
t.Fatal("payload mismatch")
|
|
||||||
}
|
|
||||||
t.Logf("%d bytes took %v (%.1f KiB/s)", len(payload), elapsed,
|
|
||||||
float64(len(payload))/1024/elapsed.Seconds())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLiveDriveParallelPut(t *testing.T) {
|
|
||||||
config := liveDriveConfig(t)
|
|
||||||
storage, err := newDriveStorage(nil, config)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newDriveStorage: %v", err)
|
|
||||||
}
|
|
||||||
defer storage.Close()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
defer storage.Delete(ctx, "streams/benchtest")
|
|
||||||
|
|
||||||
chunk := make([]byte, 256*1024)
|
|
||||||
if _, err := rand.Read(chunk); err != nil {
|
|
||||||
t.Fatalf("rand: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := storage.Put(ctx, "streams/benchtest/warmup", chunk); err != nil {
|
|
||||||
t.Fatalf("warmup: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
start := time.Now()
|
|
||||||
for i := 0; i < 4; i++ {
|
|
||||||
if err := storage.Put(ctx, fmt.Sprintf("streams/benchtest/seq%d", i), chunk); err != nil {
|
|
||||||
t.Fatalf("sequential put: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
sequential := time.Since(start)
|
|
||||||
t.Logf("4 sequential puts of 256 KiB: %v (%.1f KiB/s)",
|
|
||||||
sequential, float64(4*len(chunk))/1024/sequential.Seconds())
|
|
||||||
|
|
||||||
start = time.Now()
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
failures := make([]error, 8)
|
|
||||||
for i := 0; i < 8; i++ {
|
|
||||||
wg.Add(1)
|
|
||||||
go func(i int) {
|
|
||||||
defer wg.Done()
|
|
||||||
failures[i] = storage.Put(ctx, fmt.Sprintf("streams/benchtest/par%d", i), chunk)
|
|
||||||
}(i)
|
|
||||||
}
|
|
||||||
wg.Wait()
|
|
||||||
parallel := time.Since(start)
|
|
||||||
for _, err := range failures {
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("parallel put: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
t.Logf("8 parallel puts of 256 KiB: %v (%.1f KiB/s)",
|
|
||||||
parallel, float64(8*len(chunk))/1024/parallel.Seconds())
|
|
||||||
|
|
||||||
start = time.Now()
|
|
||||||
names, err := storage.List(ctx, "streams/benchtest")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List: %v", err)
|
|
||||||
}
|
|
||||||
t.Logf("List of %d objects took %v", len(names), time.Since(start))
|
|
||||||
|
|
||||||
start = time.Now()
|
|
||||||
wg = sync.WaitGroup{}
|
|
||||||
for i := 0; i < 8; i++ {
|
|
||||||
wg.Add(1)
|
|
||||||
go func(i int) {
|
|
||||||
defer wg.Done()
|
|
||||||
storage.Get(ctx, fmt.Sprintf("streams/benchtest/par%d", i))
|
|
||||||
}(i)
|
|
||||||
}
|
|
||||||
wg.Wait()
|
|
||||||
download := time.Since(start)
|
|
||||||
t.Logf("8 parallel gets of 256 KiB: %v (%.1f KiB/s)",
|
|
||||||
download, float64(8*len(chunk))/1024/download.Seconds())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLiveDriveSegmentSweep(t *testing.T) {
|
|
||||||
config := liveDriveConfig(t)
|
|
||||||
storage, err := newDriveStorage(nil, config)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newDriveStorage: %v", err)
|
|
||||||
}
|
|
||||||
defer storage.Close()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
defer storage.Delete(ctx, "streams/sweeptest")
|
|
||||||
|
|
||||||
const total = 1024 * 1024
|
|
||||||
for _, size := range []int{64 * 1024, 128 * 1024, 256 * 1024, 512 * 1024} {
|
|
||||||
chunk := make([]byte, size)
|
|
||||||
if _, err := rand.Read(chunk); err != nil {
|
|
||||||
t.Fatalf("rand: %v", err)
|
|
||||||
}
|
|
||||||
count := total / size
|
|
||||||
|
|
||||||
start := time.Now()
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
for i := 0; i < count; i++ {
|
|
||||||
wg.Add(1)
|
|
||||||
go func(i int) {
|
|
||||||
defer wg.Done()
|
|
||||||
storage.Put(ctx, fmt.Sprintf("streams/sweeptest/s%d-%d", size, i), chunk)
|
|
||||||
}(i)
|
|
||||||
}
|
|
||||||
wg.Wait()
|
|
||||||
elapsed := time.Since(start)
|
|
||||||
|
|
||||||
t.Logf("%4d KiB x %2d = 1 MiB in %8v -> %6.1f KiB/s",
|
|
||||||
size/1024, count, elapsed.Round(time.Millisecond),
|
|
||||||
float64(total)/1024/elapsed.Seconds())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLiveDriveListLag(t *testing.T) {
|
|
||||||
config := liveDriveConfig(t)
|
|
||||||
storage, err := newDriveStorage(nil, config)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newDriveStorage: %v", err)
|
|
||||||
}
|
|
||||||
defer storage.Close()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
defer storage.Delete(ctx, "streams/lagtest")
|
|
||||||
|
|
||||||
const rounds = 6
|
|
||||||
var worst time.Duration
|
|
||||||
|
|
||||||
for i := 0; i < rounds; i++ {
|
|
||||||
name := fmt.Sprintf("streams/lagtest/round%d/000000000.seg", i)
|
|
||||||
if err := storage.Put(ctx, name, []byte("probe")); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
start := time.Now()
|
|
||||||
var lag time.Duration
|
|
||||||
for {
|
|
||||||
names, err := storage.List(ctx, fmt.Sprintf("streams/lagtest/round%d", i))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List: %v", err)
|
|
||||||
}
|
|
||||||
if len(names) == 1 {
|
|
||||||
lag = time.Since(start)
|
|
||||||
break
|
|
||||||
}
|
|
||||||
if time.Since(start) > 30*time.Second {
|
|
||||||
t.Fatalf("round %d: the object never showed up in a listing", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if lag > worst {
|
|
||||||
worst = lag
|
|
||||||
}
|
|
||||||
t.Logf("round %d: the object became listable after %v", i, lag.Round(time.Millisecond))
|
|
||||||
}
|
|
||||||
t.Logf("worst listing lag: %v", worst.Round(time.Millisecond))
|
|
||||||
}
|
|
||||||
@@ -1,690 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"mime"
|
|
||||||
"mime/multipart"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"net/url"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
)
|
|
||||||
|
|
||||||
type fakeFile struct {
|
|
||||||
id string
|
|
||||||
name string
|
|
||||||
data []byte
|
|
||||||
description string
|
|
||||||
}
|
|
||||||
|
|
||||||
type fakeDrive struct {
|
|
||||||
server *httptest.Server
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
files map[string]*fakeFile
|
|
||||||
nextID int
|
|
||||||
failOnce map[string]bool
|
|
||||||
failStatus map[string]int
|
|
||||||
failBody map[string]string
|
|
||||||
tokens int
|
|
||||||
hosts map[string]bool
|
|
||||||
}
|
|
||||||
|
|
||||||
func newFakeDrive(t *testing.T) *fakeDrive {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
drive := &fakeDrive{
|
|
||||||
files: make(map[string]*fakeFile),
|
|
||||||
failOnce: make(map[string]bool),
|
|
||||||
failStatus: make(map[string]int),
|
|
||||||
failBody: make(map[string]string),
|
|
||||||
hosts: make(map[string]bool),
|
|
||||||
}
|
|
||||||
mux := http.NewServeMux()
|
|
||||||
mux.Handle("/", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
http.NotFound(w, r)
|
|
||||||
}))
|
|
||||||
record := func(next http.HandlerFunc) http.HandlerFunc {
|
|
||||||
return func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
drive.mu.Lock()
|
|
||||||
drive.hosts[r.Host] = true
|
|
||||||
drive.mu.Unlock()
|
|
||||||
next(w, r)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
mux.HandleFunc("/token", record(drive.handleToken))
|
|
||||||
mux.HandleFunc("/upload", record(drive.handleUpload))
|
|
||||||
mux.HandleFunc("/files", record(drive.handleFiles))
|
|
||||||
mux.HandleFunc("/files/", record(drive.handleFile))
|
|
||||||
drive.server = httptest.NewServer(mux)
|
|
||||||
|
|
||||||
resetSharedStorage()
|
|
||||||
|
|
||||||
previous := []string{driveTokenURL, driveFilesURL, driveUploadURL}
|
|
||||||
previousBackoff := driveInitialBackoff
|
|
||||||
driveTokenURL = drive.server.URL + "/token"
|
|
||||||
driveFilesURL = drive.server.URL + "/files"
|
|
||||||
driveUploadURL = drive.server.URL + "/upload"
|
|
||||||
driveInitialBackoff = 5 * time.Millisecond
|
|
||||||
|
|
||||||
t.Cleanup(func() {
|
|
||||||
driveTokenURL, driveFilesURL, driveUploadURL = previous[0], previous[1], previous[2]
|
|
||||||
driveInitialBackoff = previousBackoff
|
|
||||||
drive.server.Close()
|
|
||||||
resetSharedStorage()
|
|
||||||
})
|
|
||||||
return drive
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) handleToken(w http.ResponseWriter, r *http.Request) {
|
|
||||||
d.mu.Lock()
|
|
||||||
d.tokens++
|
|
||||||
d.mu.Unlock()
|
|
||||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
|
||||||
"access_token": "fake-token",
|
|
||||||
"expires_in": 3600,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) handleUpload(w http.ResponseWriter, r *http.Request) {
|
|
||||||
if d.shouldFail("upload") {
|
|
||||||
if status, body := d.failure("upload"); status != 0 {
|
|
||||||
w.WriteHeader(status)
|
|
||||||
w.Write([]byte(body))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
w.WriteHeader(http.StatusTooManyRequests)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
_, params, err := mime.ParseMediaType(r.Header.Get("Content-Type"))
|
|
||||||
if err != nil {
|
|
||||||
w.WriteHeader(http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
reader := multipart.NewReader(r.Body, params["boundary"])
|
|
||||||
|
|
||||||
metaPart, err := reader.NextPart()
|
|
||||||
if err != nil {
|
|
||||||
w.WriteHeader(http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
var metadata struct {
|
|
||||||
Name string `json:"name"`
|
|
||||||
}
|
|
||||||
if err := json.NewDecoder(metaPart).Decode(&metadata); err != nil {
|
|
||||||
w.WriteHeader(http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
var data []byte
|
|
||||||
if dataPart, err := reader.NextPart(); err == nil {
|
|
||||||
data, _ = io.ReadAll(dataPart)
|
|
||||||
}
|
|
||||||
|
|
||||||
d.mu.Lock()
|
|
||||||
d.nextID++
|
|
||||||
id := fmt.Sprintf("id-%d", d.nextID)
|
|
||||||
d.files[id] = &fakeFile{id: id, name: metadata.Name, data: data}
|
|
||||||
d.mu.Unlock()
|
|
||||||
|
|
||||||
json.NewEncoder(w).Encode(map[string]string{"id": id, "name": metadata.Name})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) handleFiles(w http.ResponseWriter, r *http.Request) {
|
|
||||||
if r.Method == http.MethodPost {
|
|
||||||
d.handleCreate(w, r)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
d.handleList(w, r)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) handleCreate(w http.ResponseWriter, r *http.Request) {
|
|
||||||
if d.shouldFail("upload") {
|
|
||||||
if status, body := d.failure("upload"); status != 0 {
|
|
||||||
w.WriteHeader(status)
|
|
||||||
w.Write([]byte(body))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
w.WriteHeader(http.StatusTooManyRequests)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
var meta struct {
|
|
||||||
Name string `json:"name"`
|
|
||||||
Description string `json:"description"`
|
|
||||||
}
|
|
||||||
if err := json.NewDecoder(r.Body).Decode(&meta); err != nil {
|
|
||||||
w.WriteHeader(http.StatusBadRequest)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
d.mu.Lock()
|
|
||||||
d.nextID++
|
|
||||||
id := fmt.Sprintf("id-%d", d.nextID)
|
|
||||||
d.files[id] = &fakeFile{id: id, name: meta.Name, description: meta.Description}
|
|
||||||
d.mu.Unlock()
|
|
||||||
|
|
||||||
json.NewEncoder(w).Encode(map[string]string{"id": id, "name": meta.Name})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) handleList(w http.ResponseWriter, r *http.Request) {
|
|
||||||
if d.shouldFail("list") {
|
|
||||||
w.WriteHeader(http.StatusServiceUnavailable)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
query := r.URL.Query().Get("q")
|
|
||||||
exact, prefix := parseFakeQuery(query)
|
|
||||||
|
|
||||||
type entry struct {
|
|
||||||
ID string `json:"id"`
|
|
||||||
Name string `json:"name"`
|
|
||||||
Description string `json:"description,omitempty"`
|
|
||||||
}
|
|
||||||
result := struct {
|
|
||||||
Files []entry `json:"files"`
|
|
||||||
}{}
|
|
||||||
|
|
||||||
d.mu.Lock()
|
|
||||||
for _, file := range d.files {
|
|
||||||
match := false
|
|
||||||
switch {
|
|
||||||
case exact != "":
|
|
||||||
match = file.name == exact
|
|
||||||
case prefix != "":
|
|
||||||
match = strings.HasPrefix(file.name, prefix)
|
|
||||||
}
|
|
||||||
if match {
|
|
||||||
result.Files = append(result.Files, entry{
|
|
||||||
ID: file.id, Name: file.name, Description: file.description,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
d.mu.Unlock()
|
|
||||||
|
|
||||||
json.NewEncoder(w).Encode(result)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) handleFile(w http.ResponseWriter, r *http.Request) {
|
|
||||||
id := strings.TrimPrefix(r.URL.Path, "/files/")
|
|
||||||
|
|
||||||
d.mu.Lock()
|
|
||||||
file, ok := d.files[id]
|
|
||||||
if ok && r.Method == http.MethodDelete {
|
|
||||||
delete(d.files, id)
|
|
||||||
}
|
|
||||||
d.mu.Unlock()
|
|
||||||
|
|
||||||
if !ok {
|
|
||||||
w.WriteHeader(http.StatusNotFound)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if r.Method == http.MethodDelete {
|
|
||||||
w.WriteHeader(http.StatusNoContent)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if strings.Contains(r.URL.RawQuery, "fields=description") {
|
|
||||||
json.NewEncoder(w).Encode(map[string]string{"description": file.description})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
w.Write(file.data)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) shouldFail(kind string) bool {
|
|
||||||
d.mu.Lock()
|
|
||||||
defer d.mu.Unlock()
|
|
||||||
if d.failOnce[kind] {
|
|
||||||
d.failOnce[kind] = false
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) failOnceWith(kind string, status int, body string) {
|
|
||||||
d.mu.Lock()
|
|
||||||
d.failOnce[kind] = true
|
|
||||||
d.failStatus[kind] = status
|
|
||||||
d.failBody[kind] = body
|
|
||||||
d.mu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) failure(kind string) (int, string) {
|
|
||||||
d.mu.Lock()
|
|
||||||
defer d.mu.Unlock()
|
|
||||||
if status, ok := d.failStatus[kind]; ok {
|
|
||||||
return status, d.failBody[kind]
|
|
||||||
}
|
|
||||||
return 0, ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) failNext(kind string) {
|
|
||||||
d.mu.Lock()
|
|
||||||
d.failOnce[kind] = true
|
|
||||||
d.mu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) seenHosts() []string {
|
|
||||||
d.mu.Lock()
|
|
||||||
defer d.mu.Unlock()
|
|
||||||
out := make([]string, 0, len(d.hosts))
|
|
||||||
for h := range d.hosts {
|
|
||||||
out = append(out, h)
|
|
||||||
}
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *fakeDrive) count() int {
|
|
||||||
d.mu.Lock()
|
|
||||||
defer d.mu.Unlock()
|
|
||||||
return len(d.files)
|
|
||||||
}
|
|
||||||
|
|
||||||
func parseFakeQuery(query string) (exact, prefix string) {
|
|
||||||
if value, ok := cutQuoted(query, "name = '"); ok {
|
|
||||||
return value, ""
|
|
||||||
}
|
|
||||||
if value, ok := cutQuoted(query, "name contains '"); ok {
|
|
||||||
return "", value
|
|
||||||
}
|
|
||||||
return "", ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func cutQuoted(query, marker string) (string, bool) {
|
|
||||||
start := strings.Index(query, marker)
|
|
||||||
if start < 0 {
|
|
||||||
return "", false
|
|
||||||
}
|
|
||||||
rest := query[start+len(marker):]
|
|
||||||
end := strings.Index(rest, "'")
|
|
||||||
if end < 0 {
|
|
||||||
return "", false
|
|
||||||
}
|
|
||||||
return rest[:end], true
|
|
||||||
}
|
|
||||||
|
|
||||||
func driveSettings() *internet.MemoryStreamConfig {
|
|
||||||
return &internet.MemoryStreamConfig{
|
|
||||||
ProtocolName: protocolName,
|
|
||||||
ProtocolSettings: &Config{
|
|
||||||
RemoteFolder: "folder-id",
|
|
||||||
Service: "Google Drive",
|
|
||||||
Secrets: []string{"client", "secret", "refresh"},
|
|
||||||
FlushIntervalMs: 5,
|
|
||||||
PollIntervalMs: 5,
|
|
||||||
MaxPollIntervalMs: 20,
|
|
||||||
SessionTtlSeconds: 1,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newDriveBackend(t *testing.T) *driveStorage {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
storage, err := newDriveStorage(driveSettings(), driveSettings().ProtocolSettings.(*Config))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newDriveStorage: %v", err)
|
|
||||||
}
|
|
||||||
return storage
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveSecrets(t *testing.T) {
|
|
||||||
if _, err := newDriveStorage(nil, &Config{RemoteFolder: "f", Secrets: []string{"a", "b"}}); err == nil {
|
|
||||||
t.Fatal("accepted two secrets")
|
|
||||||
}
|
|
||||||
if _, err := newDriveStorage(nil, &Config{Secrets: []string{"a", "b", "c"}}); err == nil {
|
|
||||||
t.Fatal("accepted an empty remoteFolder")
|
|
||||||
}
|
|
||||||
if _, err := newDriveStorage(nil, &Config{RemoteFolder: "f", Secrets: []string{"a", "", "c"}}); err == nil {
|
|
||||||
t.Fatal("accepted an empty secret")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveRoundTrip(t *testing.T) {
|
|
||||||
drive := newFakeDrive(t)
|
|
||||||
storage := newDriveBackend(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
if err := storage.Put(ctx, "streams/abc/c2s/000000000.seg", []byte("hello")); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
data, err := storage.Get(ctx, "streams/abc/c2s/000000000.seg")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Get: %v", err)
|
|
||||||
}
|
|
||||||
if string(data) != "hello" {
|
|
||||||
t.Fatalf("Get returned %q, want %q", data, "hello")
|
|
||||||
}
|
|
||||||
|
|
||||||
names, err := storage.List(ctx, "streams/abc/c2s")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List: %v", err)
|
|
||||||
}
|
|
||||||
if len(names) != 1 || names[0].Name != "000000000.seg" {
|
|
||||||
t.Fatalf("List returned %v, want 1 segment", names)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := storage.Get(ctx, "streams/abc/c2s/000000009.seg"); err != errNotFound {
|
|
||||||
t.Fatalf("Get returned %v, want errNotFound", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := storage.Delete(ctx, "streams/abc/c2s/000000000.seg"); err != nil {
|
|
||||||
t.Fatalf("Delete: %v", err)
|
|
||||||
}
|
|
||||||
if drive.count() != 0 {
|
|
||||||
t.Fatalf("fake drive still holds %d files", drive.count())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveListChildren(t *testing.T) {
|
|
||||||
newFakeDrive(t)
|
|
||||||
storage := newDriveBackend(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
for _, name := range []string{
|
|
||||||
"streams/one/c2s/000000000.seg",
|
|
||||||
"streams/one/s2c/000000000.seg",
|
|
||||||
"streams/two/c2s/000000000.seg",
|
|
||||||
} {
|
|
||||||
if err := storage.Put(ctx, name, []byte("x")); err != nil {
|
|
||||||
t.Fatalf("Put %s: %v", name, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
names, err := storage.List(ctx, "streams")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List: %v", err)
|
|
||||||
}
|
|
||||||
if len(names) != 2 {
|
|
||||||
t.Fatalf("List returned %v, want 2 sessions", names)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveDeleteSession(t *testing.T) {
|
|
||||||
drive := newFakeDrive(t)
|
|
||||||
storage := newDriveBackend(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
for _, name := range []string{
|
|
||||||
"streams/one/c2s/000000000.seg",
|
|
||||||
"streams/one/c2s/000000001.end",
|
|
||||||
"streams/one/s2c/000000000.seg",
|
|
||||||
"streams/two/c2s/000000000.seg",
|
|
||||||
} {
|
|
||||||
if err := storage.Put(ctx, name, []byte("x")); err != nil {
|
|
||||||
t.Fatalf("Put %s: %v", name, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := storage.Delete(ctx, "streams/one"); err != nil {
|
|
||||||
t.Fatalf("Delete: %v", err)
|
|
||||||
}
|
|
||||||
if drive.count() != 1 {
|
|
||||||
t.Fatalf("fake drive holds %d files, want 1", drive.count())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveRetry(t *testing.T) {
|
|
||||||
drive := newFakeDrive(t)
|
|
||||||
storage := newDriveBackend(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
drive.failNext("upload")
|
|
||||||
if err := storage.Put(ctx, "sessions/abc", nil); err != nil {
|
|
||||||
t.Fatalf("Put did not survive a 429: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
drive.failNext("list")
|
|
||||||
if _, err := storage.List(ctx, "sessions"); err != nil {
|
|
||||||
t.Fatalf("List did not survive a 503: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveTokenCache(t *testing.T) {
|
|
||||||
drive := newFakeDrive(t)
|
|
||||||
storage := newDriveBackend(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
for i := 0; i < 5; i++ {
|
|
||||||
if err := storage.Put(ctx, fmt.Sprintf("sessions/s%d", i), nil); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if drive.tokens != 1 {
|
|
||||||
t.Fatalf("token endpoint hit %d times, want 1", drive.tokens)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveTransport(t *testing.T) {
|
|
||||||
newFakeDrive(t)
|
|
||||||
|
|
||||||
client, server, cleanup := pairWith(t, driveSettings())
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
if _, err := client.Write([]byte("ping")); err != nil {
|
|
||||||
t.Fatalf("client write: %v", err)
|
|
||||||
}
|
|
||||||
expectRead(t, server, "ping")
|
|
||||||
|
|
||||||
if _, err := server.Write([]byte("pong")); err != nil {
|
|
||||||
t.Fatalf("server write: %v", err)
|
|
||||||
}
|
|
||||||
expectRead(t, client, "pong")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveLargeTransfer(t *testing.T) {
|
|
||||||
newFakeDrive(t)
|
|
||||||
|
|
||||||
client, server, cleanup := pairWith(t, driveSettings())
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
payload := make([]byte, 300000)
|
|
||||||
for i := range payload {
|
|
||||||
payload[i] = byte(i % 251)
|
|
||||||
}
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
client.Write(payload)
|
|
||||||
}()
|
|
||||||
|
|
||||||
if err := server.SetReadDeadline(time.Now().Add(60 * time.Second)); err != nil {
|
|
||||||
t.Fatalf("SetReadDeadline: %v", err)
|
|
||||||
}
|
|
||||||
got := make([]byte, len(payload))
|
|
||||||
if _, err := io.ReadFull(server, got); err != nil {
|
|
||||||
t.Fatalf("ReadFull: %v", err)
|
|
||||||
}
|
|
||||||
for i := range got {
|
|
||||||
if got[i] != payload[i] {
|
|
||||||
t.Fatalf("payload mismatch at byte %d", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRateLimited(t *testing.T) {
|
|
||||||
limited := []string{
|
|
||||||
`{"error":{"code":403,"errors":[{"reason":"userRateLimitExceeded"}]}}`,
|
|
||||||
`{"error":{"code":403,"errors":[{"reason":"rateLimitExceeded"}]}}`,
|
|
||||||
`{"error":{"code":403,"errors":[{"reason":"sharingRateLimitExceeded"}]}}`,
|
|
||||||
`{"error":{"status":"RESOURCE_EXHAUSTED"}}`,
|
|
||||||
}
|
|
||||||
for _, payload := range limited {
|
|
||||||
if !rateLimited([]byte(payload)) {
|
|
||||||
t.Fatalf("rateLimited missed %s", payload)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
permanent := []string{
|
|
||||||
`{"error":{"code":403,"errors":[{"reason":"insufficientFilePermissions"}]}}`,
|
|
||||||
`{"error":{"code":403,"errors":[{"reason":"storageQuotaExceeded"}]}}`,
|
|
||||||
`not json at all`,
|
|
||||||
}
|
|
||||||
for _, payload := range permanent {
|
|
||||||
if rateLimited([]byte(payload)) {
|
|
||||||
t.Fatalf("rateLimited treated %s as temporary", payload)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveRetryRateLimit(t *testing.T) {
|
|
||||||
drive := newFakeDrive(t)
|
|
||||||
storage := newDriveBackend(t)
|
|
||||||
|
|
||||||
drive.failOnceWith("upload", http.StatusForbidden,
|
|
||||||
`{"error":{"code":403,"errors":[{"reason":"userRateLimitExceeded"}]}}`)
|
|
||||||
|
|
||||||
if err := storage.Put(context.Background(), "sessions/abc", nil); err != nil {
|
|
||||||
t.Fatalf("Put did not survive a 403: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveSharedClient(t *testing.T) {
|
|
||||||
newFakeDrive(t)
|
|
||||||
|
|
||||||
first, err := newStorage(driveSettings())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newStorage: %v", err)
|
|
||||||
}
|
|
||||||
second, err := newStorage(driveSettings())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newStorage: %v", err)
|
|
||||||
}
|
|
||||||
if first != second {
|
|
||||||
t.Fatal("same settings did not share one storage")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveInlineListing(t *testing.T) {
|
|
||||||
drive := newFakeDrive(t)
|
|
||||||
storage := newDriveBackend(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
payload := []byte("small enough to ride along with the listing")
|
|
||||||
if err := storage.Put(ctx, "streams/abc/c2s/000000000.seg", payload); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
entries, err := storage.List(ctx, "streams/abc/c2s")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List: %v", err)
|
|
||||||
}
|
|
||||||
if len(entries) != 1 {
|
|
||||||
t.Fatalf("List returned %v, want 1 entry", entries)
|
|
||||||
}
|
|
||||||
if string(entries[0].Inline) != string(payload) {
|
|
||||||
t.Fatalf("listing carried %q, want %q", entries[0].Inline, payload)
|
|
||||||
}
|
|
||||||
|
|
||||||
d := drive
|
|
||||||
d.mu.Lock()
|
|
||||||
var stored *fakeFile
|
|
||||||
for _, f := range d.files {
|
|
||||||
stored = f
|
|
||||||
}
|
|
||||||
d.mu.Unlock()
|
|
||||||
if len(stored.data) != 0 {
|
|
||||||
t.Fatal("small payload was uploaded as content")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveLargeContent(t *testing.T) {
|
|
||||||
drive := newFakeDrive(t)
|
|
||||||
storage := newDriveBackend(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
payload := make([]byte, driveInlineLimit+1)
|
|
||||||
for i := range payload {
|
|
||||||
payload[i] = byte(i)
|
|
||||||
}
|
|
||||||
if err := storage.Put(ctx, "streams/abc/c2s/000000000.seg", payload); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
entries, err := storage.List(ctx, "streams/abc/c2s")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List: %v", err)
|
|
||||||
}
|
|
||||||
if len(entries) != 1 || entries[0].Inline != nil {
|
|
||||||
t.Fatalf("large payload was inlined, got %v", entries)
|
|
||||||
}
|
|
||||||
|
|
||||||
got, err := storage.Get(ctx, "streams/abc/c2s/000000000.seg")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Get: %v", err)
|
|
||||||
}
|
|
||||||
if len(got) != len(payload) {
|
|
||||||
t.Fatalf("Get returned %d bytes, want %d", len(got), len(payload))
|
|
||||||
}
|
|
||||||
_ = drive
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveInlineGet(t *testing.T) {
|
|
||||||
newFakeDrive(t)
|
|
||||||
storage := newDriveBackend(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
payload := []byte("only in the description")
|
|
||||||
if err := storage.Put(ctx, "sessions/abc", payload); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
got, err := storage.Get(ctx, "sessions/abc")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Get: %v", err)
|
|
||||||
}
|
|
||||||
if string(got) != string(payload) {
|
|
||||||
t.Fatalf("Get returned %q, want %q", got, payload)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriveFronting(t *testing.T) {
|
|
||||||
drive := newFakeDrive(t)
|
|
||||||
|
|
||||||
fake, err := url.Parse(drive.server.URL)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("parse: %v", err)
|
|
||||||
}
|
|
||||||
port, err := net.PortFromString(fake.Port())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("port: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
driveTokenURL = "http://www.googleapis.com/token"
|
|
||||||
driveFilesURL = "http://www.googleapis.com/files"
|
|
||||||
driveUploadURL = "http://www.googleapis.com/upload"
|
|
||||||
|
|
||||||
settings := driveSettings()
|
|
||||||
settings.Destination = &net.Destination{
|
|
||||||
Address: net.ParseAddress(fake.Hostname()),
|
|
||||||
Port: port,
|
|
||||||
Network: net.Network_TCP,
|
|
||||||
}
|
|
||||||
|
|
||||||
storage, err := newDriveStorage(settings, settings.ProtocolSettings.(*Config))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newDriveStorage: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := storage.Put(context.Background(), "sessions/fronted", []byte("x")); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
if drive.count() != 1 {
|
|
||||||
t.Fatalf("fake drive holds %d files, want 1", drive.count())
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, host := range drive.seenHosts() {
|
|
||||||
if host != "www.googleapis.com" {
|
|
||||||
t.Fatalf("inner host was %q, want www.googleapis.com regardless of address", host)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,39 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/json"
|
|
||||||
"strings"
|
|
||||||
)
|
|
||||||
|
|
||||||
func jsonWalk(payload []byte, path string) interface{} {
|
|
||||||
var root interface{}
|
|
||||||
if json.Unmarshal(payload, &root) != nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
node := root
|
|
||||||
for _, key := range strings.Split(path, ".") {
|
|
||||||
obj, ok := node.(map[string]interface{})
|
|
||||||
if !ok {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
node, ok = obj[key]
|
|
||||||
if !ok {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return node
|
|
||||||
}
|
|
||||||
|
|
||||||
func jsonString(payload []byte, path string) string {
|
|
||||||
if s, ok := jsonWalk(payload, path).(string); ok {
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
func jsonNumber(payload []byte, path string) int64 {
|
|
||||||
if f, ok := jsonWalk(payload, path).(float64); ok {
|
|
||||||
return int64(f)
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
@@ -1,120 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"os"
|
|
||||||
"path"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
)
|
|
||||||
|
|
||||||
const tempPrefix = ".xdrive-tmp-"
|
|
||||||
|
|
||||||
type localStorage struct {
|
|
||||||
root string
|
|
||||||
}
|
|
||||||
|
|
||||||
func newLocalStorage(root string) (*localStorage, error) {
|
|
||||||
if root == "" {
|
|
||||||
return nil, errors.New(`empty "remoteFolder"`)
|
|
||||||
}
|
|
||||||
if err := os.MkdirAll(root, 0o700); err != nil {
|
|
||||||
return nil, errors.New("failed to create remote folder").Base(err)
|
|
||||||
}
|
|
||||||
return &localStorage{root: root}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *localStorage) resolve(name string) (string, error) {
|
|
||||||
clean := path.Clean("/" + name)
|
|
||||||
if clean == "/" {
|
|
||||||
return "", errors.New("invalid object name: ", name)
|
|
||||||
}
|
|
||||||
if strings.HasPrefix(path.Base(clean), tempPrefix) {
|
|
||||||
return "", errors.New("reserved object name: ", name)
|
|
||||||
}
|
|
||||||
return filepath.Join(s.root, filepath.FromSlash(clean[1:])), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *localStorage) Put(ctx context.Context, name string, data []byte) error {
|
|
||||||
full, err := s.resolve(name)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
dir := filepath.Dir(full)
|
|
||||||
if err := os.MkdirAll(dir, 0o700); err != nil {
|
|
||||||
return errors.New("failed to create folder ", dir).Base(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
tmp, err := os.CreateTemp(dir, tempPrefix+"*")
|
|
||||||
if err != nil {
|
|
||||||
return errors.New("failed to create temp file in ", dir).Base(err)
|
|
||||||
}
|
|
||||||
tmpName := tmp.Name()
|
|
||||||
defer os.Remove(tmpName)
|
|
||||||
|
|
||||||
if _, err := tmp.Write(data); err != nil {
|
|
||||||
tmp.Close()
|
|
||||||
return errors.New("failed to write ", name).Base(err)
|
|
||||||
}
|
|
||||||
if err := tmp.Close(); err != nil {
|
|
||||||
return errors.New("failed to close ", name).Base(err)
|
|
||||||
}
|
|
||||||
if err := os.Rename(tmpName, full); err != nil {
|
|
||||||
return errors.New("failed to commit ", name).Base(err)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *localStorage) Get(ctx context.Context, name string) ([]byte, error) {
|
|
||||||
full, err := s.resolve(name)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
data, err := os.ReadFile(full)
|
|
||||||
if err != nil {
|
|
||||||
if os.IsNotExist(err) {
|
|
||||||
return nil, errNotFound
|
|
||||||
}
|
|
||||||
return nil, errors.New("failed to read ", name).Base(err)
|
|
||||||
}
|
|
||||||
return data, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *localStorage) Delete(ctx context.Context, name string) error {
|
|
||||||
full, err := s.resolve(name)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if err := os.RemoveAll(full); err != nil {
|
|
||||||
return errors.New("failed to delete ", name).Base(err)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *localStorage) List(ctx context.Context, prefix string) ([]Entry, error) {
|
|
||||||
full, err := s.resolve(prefix)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
entries, err := os.ReadDir(full)
|
|
||||||
if err != nil {
|
|
||||||
if os.IsNotExist(err) {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
return nil, errors.New("failed to list ", prefix).Base(err)
|
|
||||||
}
|
|
||||||
found := make([]Entry, 0, len(entries))
|
|
||||||
for _, entry := range entries {
|
|
||||||
if strings.HasPrefix(entry.Name(), tempPrefix) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
found = append(found, Entry{Name: entry.Name()})
|
|
||||||
}
|
|
||||||
return found, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *localStorage) Close() error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,69 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import "time"
|
|
||||||
|
|
||||||
const (
|
|
||||||
defaultSegmentBytes = 512 * 1024
|
|
||||||
defaultFlushInterval = 20 * time.Millisecond
|
|
||||||
defaultMinPollInterval = 50 * time.Millisecond
|
|
||||||
defaultMaxPollInterval = 500 * time.Millisecond
|
|
||||||
defaultEagerWindow = 2 * time.Second
|
|
||||||
defaultHoleTimeout = 30 * time.Second
|
|
||||||
defaultSessionTTL = 5 * time.Minute
|
|
||||||
defaultConcurrency = 8
|
|
||||||
|
|
||||||
maxSegmentBytes = 16 * 1024 * 1024
|
|
||||||
maxConcurrency = 64
|
|
||||||
)
|
|
||||||
|
|
||||||
type params struct {
|
|
||||||
segmentBytes int
|
|
||||||
flushInterval time.Duration
|
|
||||||
minPollInterval time.Duration
|
|
||||||
maxPollInterval time.Duration
|
|
||||||
eagerWindow time.Duration
|
|
||||||
holeTimeout time.Duration
|
|
||||||
sessionTTL time.Duration
|
|
||||||
concurrency int
|
|
||||||
}
|
|
||||||
|
|
||||||
func millis(value uint32, fallback time.Duration) time.Duration {
|
|
||||||
if value == 0 {
|
|
||||||
return fallback
|
|
||||||
}
|
|
||||||
return time.Duration(value) * time.Millisecond
|
|
||||||
}
|
|
||||||
|
|
||||||
func seconds(value uint32, fallback time.Duration) time.Duration {
|
|
||||||
if value == 0 {
|
|
||||||
return fallback
|
|
||||||
}
|
|
||||||
return time.Duration(value) * time.Second
|
|
||||||
}
|
|
||||||
|
|
||||||
func capped(value uint32, fallback, limit int) int {
|
|
||||||
if value == 0 {
|
|
||||||
return fallback
|
|
||||||
}
|
|
||||||
if int(value) > limit {
|
|
||||||
return limit
|
|
||||||
}
|
|
||||||
return int(value)
|
|
||||||
}
|
|
||||||
|
|
||||||
func paramsFromConfig(c *Config) params {
|
|
||||||
p := params{
|
|
||||||
segmentBytes: capped(c.SegmentBytes, defaultSegmentBytes, maxSegmentBytes),
|
|
||||||
flushInterval: millis(c.FlushIntervalMs, defaultFlushInterval),
|
|
||||||
minPollInterval: millis(c.PollIntervalMs, defaultMinPollInterval),
|
|
||||||
maxPollInterval: millis(c.MaxPollIntervalMs, defaultMaxPollInterval),
|
|
||||||
eagerWindow: millis(c.EagerWindowMs, defaultEagerWindow),
|
|
||||||
holeTimeout: millis(c.HoleTimeoutMs, defaultHoleTimeout),
|
|
||||||
sessionTTL: seconds(c.SessionTtlSeconds, defaultSessionTTL),
|
|
||||||
concurrency: capped(c.Concurrency, defaultConcurrency, maxConcurrency),
|
|
||||||
}
|
|
||||||
if p.maxPollInterval < p.minPollInterval {
|
|
||||||
p.maxPollInterval = p.minPollInterval
|
|
||||||
}
|
|
||||||
return p
|
|
||||||
}
|
|
||||||
@@ -1,90 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"crypto/sha256"
|
|
||||||
"encoding/hex"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
)
|
|
||||||
|
|
||||||
var errNotFound = errors.New("object not found")
|
|
||||||
|
|
||||||
type Entry struct {
|
|
||||||
Name string
|
|
||||||
Inline []byte
|
|
||||||
}
|
|
||||||
|
|
||||||
type Storage interface {
|
|
||||||
Put(ctx context.Context, name string, data []byte) error
|
|
||||||
Get(ctx context.Context, name string) ([]byte, error)
|
|
||||||
Delete(ctx context.Context, name string) error
|
|
||||||
List(ctx context.Context, prefix string) ([]Entry, error)
|
|
||||||
Close() error
|
|
||||||
}
|
|
||||||
|
|
||||||
func newStorage(streamSettings *internet.MemoryStreamConfig) (Storage, error) {
|
|
||||||
config, err := streamConfig(streamSettings)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
switch config.Service {
|
|
||||||
case "local":
|
|
||||||
return newLocalStorage(config.RemoteFolder)
|
|
||||||
case "Google Drive":
|
|
||||||
return sharedStorage(streamSettings, config, func() (Storage, error) {
|
|
||||||
return newDriveStorage(streamSettings, config)
|
|
||||||
})
|
|
||||||
case "template":
|
|
||||||
return sharedStorage(streamSettings, config, func() (Storage, error) {
|
|
||||||
return newTemplateStorage(streamSettings, config)
|
|
||||||
})
|
|
||||||
default:
|
|
||||||
return nil, errors.New("unsupported service: ", config.Service)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
|
||||||
sharedMu sync.Mutex
|
|
||||||
shared = make(map[string]Storage)
|
|
||||||
)
|
|
||||||
|
|
||||||
func shareKey(streamSettings *internet.MemoryStreamConfig, config *Config) string {
|
|
||||||
parts := []string{config.Service, config.RemoteFolder}
|
|
||||||
parts = append(parts, config.Secrets...)
|
|
||||||
if streamSettings != nil {
|
|
||||||
parts = append(parts, streamSettings.SecurityType)
|
|
||||||
if streamSettings.Destination != nil {
|
|
||||||
parts = append(parts, streamSettings.Destination.NetAddr())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
sum := sha256.Sum256([]byte(strings.Join(parts, "\x00")))
|
|
||||||
return hex.EncodeToString(sum[:])
|
|
||||||
}
|
|
||||||
|
|
||||||
func resetSharedStorage() {
|
|
||||||
sharedMu.Lock()
|
|
||||||
shared = make(map[string]Storage)
|
|
||||||
sharedMu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
func sharedStorage(streamSettings *internet.MemoryStreamConfig, config *Config, build func() (Storage, error)) (Storage, error) {
|
|
||||||
key := shareKey(streamSettings, config)
|
|
||||||
|
|
||||||
sharedMu.Lock()
|
|
||||||
defer sharedMu.Unlock()
|
|
||||||
|
|
||||||
if storage, ok := shared[key]; ok {
|
|
||||||
return storage, nil
|
|
||||||
}
|
|
||||||
storage, err := build()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
shared[key] = storage
|
|
||||||
return storage, nil
|
|
||||||
}
|
|
||||||
@@ -1,445 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"encoding/base64"
|
|
||||||
"encoding/json"
|
|
||||||
"io"
|
|
||||||
"net/http"
|
|
||||||
"regexp"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
)
|
|
||||||
|
|
||||||
type authTemplate struct {
|
|
||||||
Type string `json:"type"`
|
|
||||||
Header map[string]string `json:"header"`
|
|
||||||
Username string `json:"username"`
|
|
||||||
Password string `json:"password"`
|
|
||||||
TokenURL string `json:"tokenUrl"`
|
|
||||||
Form map[string]string `json:"form"`
|
|
||||||
TokenPath string `json:"tokenPath"`
|
|
||||||
ExpiryPath string `json:"expiryPath"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type opTemplate struct {
|
|
||||||
Method string `json:"method"`
|
|
||||||
URL string `json:"url"`
|
|
||||||
Headers map[string]string `json:"headers"`
|
|
||||||
Body string `json:"body"`
|
|
||||||
NamesRegex string `json:"namesRegex"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type retryTemplate struct {
|
|
||||||
Status []int `json:"status"`
|
|
||||||
RateReason string `json:"rateReason"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type storageTemplate struct {
|
|
||||||
Flatten bool `json:"flatten"`
|
|
||||||
Concurrency int `json:"concurrency"`
|
|
||||||
Auth authTemplate `json:"auth"`
|
|
||||||
Put opTemplate `json:"put"`
|
|
||||||
Get opTemplate `json:"get"`
|
|
||||||
Delete opTemplate `json:"delete"`
|
|
||||||
List opTemplate `json:"list"`
|
|
||||||
Retry retryTemplate `json:"retry"`
|
|
||||||
|
|
||||||
names *regexp.Regexp
|
|
||||||
}
|
|
||||||
|
|
||||||
type templateStorage struct {
|
|
||||||
tmpl *storageTemplate
|
|
||||||
client *http.Client
|
|
||||||
folder string
|
|
||||||
secrets []string
|
|
||||||
|
|
||||||
inflight chan struct{}
|
|
||||||
|
|
||||||
tokenMu sync.Mutex
|
|
||||||
token string
|
|
||||||
tokenExpiry time.Time
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTemplateStorage(streamSettings *internet.MemoryStreamConfig, config *Config) (*templateStorage, error) {
|
|
||||||
tmpl := &storageTemplate{}
|
|
||||||
if err := json.Unmarshal([]byte(config.Template), tmpl); err != nil {
|
|
||||||
return nil, errors.New("invalid template").Base(err)
|
|
||||||
}
|
|
||||||
if tmpl.Put.URL == "" || tmpl.Get.URL == "" || tmpl.List.URL == "" || tmpl.Delete.URL == "" {
|
|
||||||
return nil, errors.New("template needs put, get, list and delete operations")
|
|
||||||
}
|
|
||||||
if tmpl.List.NamesRegex == "" {
|
|
||||||
return nil, errors.New("template list needs a namesRegex")
|
|
||||||
}
|
|
||||||
re, err := regexp.Compile(tmpl.List.NamesRegex)
|
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("bad namesRegex").Base(err)
|
|
||||||
}
|
|
||||||
if re.NumSubexp() < 1 {
|
|
||||||
return nil, errors.New("namesRegex needs one capture group")
|
|
||||||
}
|
|
||||||
tmpl.names = re
|
|
||||||
|
|
||||||
conc := tmpl.Concurrency
|
|
||||||
if conc <= 0 {
|
|
||||||
conc = driveMaxInflight
|
|
||||||
}
|
|
||||||
if conc > maxTemplateConcurrency {
|
|
||||||
conc = maxTemplateConcurrency
|
|
||||||
}
|
|
||||||
|
|
||||||
return &templateStorage{
|
|
||||||
tmpl: tmpl,
|
|
||||||
client: newServiceClient(streamSettings, driveTimeout, conc),
|
|
||||||
folder: config.RemoteFolder,
|
|
||||||
secrets: config.Secrets,
|
|
||||||
inflight: make(chan struct{}, conc),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
const maxTemplateConcurrency = 256
|
|
||||||
|
|
||||||
func (s *templateStorage) baseVars() map[string]string {
|
|
||||||
vars := map[string]string{"folder": s.folder}
|
|
||||||
for i, secret := range s.secrets {
|
|
||||||
vars["secret"+itoa(i)] = secret
|
|
||||||
}
|
|
||||||
return vars
|
|
||||||
}
|
|
||||||
|
|
||||||
func subst(tmpl string, vars map[string]string) string {
|
|
||||||
if tmpl == "" || !strings.ContainsRune(tmpl, '{') {
|
|
||||||
return tmpl
|
|
||||||
}
|
|
||||||
out := tmpl
|
|
||||||
for k, v := range vars {
|
|
||||||
out = strings.ReplaceAll(out, "{"+k+"}", v)
|
|
||||||
}
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func itoa(i int) string {
|
|
||||||
if i == 0 {
|
|
||||||
return "0"
|
|
||||||
}
|
|
||||||
var b [20]byte
|
|
||||||
pos := len(b)
|
|
||||||
for i > 0 {
|
|
||||||
pos--
|
|
||||||
b[pos] = byte('0' + i%10)
|
|
||||||
i /= 10
|
|
||||||
}
|
|
||||||
return string(b[pos:])
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *templateStorage) storedName(name string) string {
|
|
||||||
if s.tmpl.Flatten {
|
|
||||||
return flatten(name)
|
|
||||||
}
|
|
||||||
return name
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *templateStorage) retryable(status int, payload []byte) bool {
|
|
||||||
for _, code := range s.tmpl.Retry.Status {
|
|
||||||
if status == code {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if status == http.StatusForbidden && s.tmpl.Retry.RateReason != "" {
|
|
||||||
if reason := jsonString(payload, s.tmpl.Retry.RateReason); reason != "" {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *templateStorage) authHeaders(ctx context.Context, vars map[string]string) (map[string]string, error) {
|
|
||||||
switch s.tmpl.Auth.Type {
|
|
||||||
case "", "none":
|
|
||||||
return nil, nil
|
|
||||||
case "static", "oauth2":
|
|
||||||
if s.tmpl.Auth.Type == "oauth2" {
|
|
||||||
token, err := s.accessToken(ctx)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
vars["token"] = token
|
|
||||||
}
|
|
||||||
headers := make(map[string]string, len(s.tmpl.Auth.Header))
|
|
||||||
for k, v := range s.tmpl.Auth.Header {
|
|
||||||
headers[k] = subst(v, vars)
|
|
||||||
}
|
|
||||||
return headers, nil
|
|
||||||
case "basic":
|
|
||||||
user := subst(s.tmpl.Auth.Username, vars)
|
|
||||||
pass := subst(s.tmpl.Auth.Password, vars)
|
|
||||||
enc := base64.StdEncoding.EncodeToString([]byte(user + ":" + pass))
|
|
||||||
return map[string]string{"Authorization": "Basic " + enc}, nil
|
|
||||||
default:
|
|
||||||
return nil, errors.New("unsupported auth type: ", s.tmpl.Auth.Type)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *templateStorage) accessToken(ctx context.Context) (string, error) {
|
|
||||||
s.tokenMu.Lock()
|
|
||||||
defer s.tokenMu.Unlock()
|
|
||||||
|
|
||||||
if s.token != "" && time.Now().Before(s.tokenExpiry) {
|
|
||||||
return s.token, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
form := make(map[string]string, len(s.tmpl.Auth.Form))
|
|
||||||
vars := s.baseVars()
|
|
||||||
values := strings.Builder{}
|
|
||||||
first := true
|
|
||||||
for k, v := range s.tmpl.Auth.Form {
|
|
||||||
form[k] = subst(v, vars)
|
|
||||||
if !first {
|
|
||||||
values.WriteByte('&')
|
|
||||||
}
|
|
||||||
first = false
|
|
||||||
values.WriteString(k)
|
|
||||||
values.WriteByte('=')
|
|
||||||
values.WriteString(form[k])
|
|
||||||
}
|
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, s.tmpl.Auth.TokenURL,
|
|
||||||
strings.NewReader(values.String()))
|
|
||||||
if err != nil {
|
|
||||||
return "", errors.New("failed to build the token request").Base(err)
|
|
||||||
}
|
|
||||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
|
||||||
|
|
||||||
resp, err := s.client.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
return "", errors.New("failed to fetch the token").Base(err)
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
payload, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
|
|
||||||
if err != nil {
|
|
||||||
return "", errors.New("failed to read the token response").Base(err)
|
|
||||||
}
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
|
||||||
return "", errors.New("the token endpoint answered ", resp.StatusCode, ": ", string(payload))
|
|
||||||
}
|
|
||||||
|
|
||||||
path := s.tmpl.Auth.TokenPath
|
|
||||||
if path == "" {
|
|
||||||
path = "access_token"
|
|
||||||
}
|
|
||||||
token := jsonString(payload, path)
|
|
||||||
if token == "" {
|
|
||||||
return "", errors.New("the token response has no token at ", path)
|
|
||||||
}
|
|
||||||
|
|
||||||
lifetime := int64(3600)
|
|
||||||
if s.tmpl.Auth.ExpiryPath != "" {
|
|
||||||
if n := jsonNumber(payload, s.tmpl.Auth.ExpiryPath); n > 0 {
|
|
||||||
lifetime = n
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if lifetime > 60 {
|
|
||||||
lifetime -= 60
|
|
||||||
}
|
|
||||||
s.token = token
|
|
||||||
s.tokenExpiry = time.Now().Add(time.Duration(lifetime) * time.Second)
|
|
||||||
return s.token, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *templateStorage) invalidateToken() {
|
|
||||||
s.tokenMu.Lock()
|
|
||||||
s.token = ""
|
|
||||||
s.tokenMu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *templateStorage) do(ctx context.Context, op *opTemplate, vars map[string]string, body []byte) (int, []byte, error) {
|
|
||||||
backoff := driveInitialBackoff
|
|
||||||
var lastErr error
|
|
||||||
|
|
||||||
for attempt := 0; attempt < driveMaxAttempts; attempt++ {
|
|
||||||
if attempt > 0 {
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
return 0, nil, ctx.Err()
|
|
||||||
case <-time.After(jitter(backoff)):
|
|
||||||
}
|
|
||||||
backoff *= 2
|
|
||||||
if backoff > driveMaxBackoff {
|
|
||||||
backoff = driveMaxBackoff
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
authHeaders, err := s.authHeaders(ctx, vars)
|
|
||||||
if err != nil {
|
|
||||||
lastErr = err
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
method := op.Method
|
|
||||||
if method == "" {
|
|
||||||
method = http.MethodGet
|
|
||||||
}
|
|
||||||
|
|
||||||
var reader io.Reader
|
|
||||||
if body != nil {
|
|
||||||
reader = bytes.NewReader(body)
|
|
||||||
}
|
|
||||||
req, err := http.NewRequestWithContext(ctx, method, subst(op.URL, vars), reader)
|
|
||||||
if err != nil {
|
|
||||||
return 0, nil, errors.New("failed to build request").Base(err)
|
|
||||||
}
|
|
||||||
for k, v := range authHeaders {
|
|
||||||
req.Header.Set(k, v)
|
|
||||||
}
|
|
||||||
for k, v := range op.Headers {
|
|
||||||
req.Header.Set(k, subst(v, vars))
|
|
||||||
}
|
|
||||||
|
|
||||||
select {
|
|
||||||
case s.inflight <- struct{}{}:
|
|
||||||
case <-ctx.Done():
|
|
||||||
return 0, nil, ctx.Err()
|
|
||||||
}
|
|
||||||
resp, err := s.client.Do(req)
|
|
||||||
<-s.inflight
|
|
||||||
if err != nil {
|
|
||||||
if ctx.Err() != nil {
|
|
||||||
return 0, nil, ctx.Err()
|
|
||||||
}
|
|
||||||
lastErr = errors.New("request failed").Base(err)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
payload, err := io.ReadAll(resp.Body)
|
|
||||||
resp.Body.Close()
|
|
||||||
if err != nil {
|
|
||||||
if ctx.Err() != nil {
|
|
||||||
return 0, nil, ctx.Err()
|
|
||||||
}
|
|
||||||
lastErr = errors.New("failed to read response").Base(err)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if resp.StatusCode == http.StatusUnauthorized && s.tmpl.Auth.Type == "oauth2" {
|
|
||||||
s.invalidateToken()
|
|
||||||
lastErr = errors.New("the service rejected the token")
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if s.retryable(resp.StatusCode, payload) {
|
|
||||||
lastErr = errors.New("the service answered ", resp.StatusCode)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
return resp.StatusCode, payload, nil
|
|
||||||
}
|
|
||||||
return 0, nil, lastErr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *templateStorage) Put(ctx context.Context, name string, data []byte) error {
|
|
||||||
vars := s.baseVars()
|
|
||||||
vars["name"] = s.storedName(name)
|
|
||||||
|
|
||||||
body := data
|
|
||||||
if s.tmpl.Put.Body != "" {
|
|
||||||
vars["data"] = base64.StdEncoding.EncodeToString(data)
|
|
||||||
body = []byte(subst(s.tmpl.Put.Body, vars))
|
|
||||||
}
|
|
||||||
|
|
||||||
status, payload, err := s.do(ctx, &s.tmpl.Put, vars, body)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if status < 200 || status >= 300 {
|
|
||||||
return errors.New("put of ", name, " answered ", status, ": ", string(payload))
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *templateStorage) Get(ctx context.Context, name string) ([]byte, error) {
|
|
||||||
vars := s.baseVars()
|
|
||||||
vars["name"] = s.storedName(name)
|
|
||||||
|
|
||||||
status, payload, err := s.do(ctx, &s.tmpl.Get, vars, nil)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
switch {
|
|
||||||
case status >= 200 && status < 300:
|
|
||||||
return payload, nil
|
|
||||||
case status == http.StatusNotFound:
|
|
||||||
return nil, errNotFound
|
|
||||||
default:
|
|
||||||
return nil, errors.New("get of ", name, " answered ", status, ": ", string(payload))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *templateStorage) Delete(ctx context.Context, name string) error {
|
|
||||||
vars := s.baseVars()
|
|
||||||
vars["name"] = s.storedName(name)
|
|
||||||
|
|
||||||
status, payload, err := s.do(ctx, &s.tmpl.Delete, vars, nil)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if status == http.StatusNotFound || (status >= 200 && status < 300) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return errors.New("delete of ", name, " answered ", status, ": ", string(payload))
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *templateStorage) List(ctx context.Context, prefix string) ([]Entry, error) {
|
|
||||||
vars := s.baseVars()
|
|
||||||
flat := s.storedName(prefix)
|
|
||||||
vars["prefix"] = flat
|
|
||||||
|
|
||||||
status, payload, err := s.do(ctx, &s.tmpl.List, vars, nil)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if status == http.StatusNotFound {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
if status < 200 || status >= 300 {
|
|
||||||
return nil, errors.New("list of ", prefix, " answered ", status, ": ", string(payload))
|
|
||||||
}
|
|
||||||
|
|
||||||
matches := s.tmpl.names.FindAllStringSubmatch(string(payload), -1)
|
|
||||||
if !s.tmpl.Flatten {
|
|
||||||
entries := make([]Entry, 0, len(matches))
|
|
||||||
for _, m := range matches {
|
|
||||||
entries = append(entries, Entry{Name: m[1]})
|
|
||||||
}
|
|
||||||
return entries, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
want := flat + flatSeparator
|
|
||||||
seen := make(map[string]bool, len(matches))
|
|
||||||
entries := make([]Entry, 0, len(matches))
|
|
||||||
for _, m := range matches {
|
|
||||||
name := m[1]
|
|
||||||
if !strings.HasPrefix(name, want) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
rest := strings.TrimPrefix(name, want)
|
|
||||||
if rest == "" {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if cut := strings.Index(rest, flatSeparator); cut >= 0 {
|
|
||||||
rest = rest[:cut]
|
|
||||||
}
|
|
||||||
if seen[rest] {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
seen[rest] = true
|
|
||||||
entries = append(entries, Entry{Name: rest})
|
|
||||||
}
|
|
||||||
return entries, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *templateStorage) Close() error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,308 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
)
|
|
||||||
|
|
||||||
type fakeStore struct {
|
|
||||||
server *httptest.Server
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
objects map[string][]byte
|
|
||||||
needAuth string
|
|
||||||
sawAuth string
|
|
||||||
tokens int
|
|
||||||
}
|
|
||||||
|
|
||||||
func newFakeStore(t *testing.T) *fakeStore {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
store := &fakeStore{objects: make(map[string][]byte)}
|
|
||||||
store.server = httptest.NewServer(http.HandlerFunc(store.handle))
|
|
||||||
t.Cleanup(store.server.Close)
|
|
||||||
return store
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *fakeStore) handle(w http.ResponseWriter, r *http.Request) {
|
|
||||||
if r.URL.Path == "/token" {
|
|
||||||
s.mu.Lock()
|
|
||||||
s.tokens++
|
|
||||||
s.mu.Unlock()
|
|
||||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
|
||||||
"access_token": "tok-fake", "expires_in": 3600,
|
|
||||||
})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if auth := r.Header.Get("Authorization"); auth != "" {
|
|
||||||
s.mu.Lock()
|
|
||||||
s.sawAuth = auth
|
|
||||||
s.mu.Unlock()
|
|
||||||
}
|
|
||||||
if s.needAuth != "" && r.Header.Get("Authorization") != s.needAuth {
|
|
||||||
w.WriteHeader(http.StatusUnauthorized)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
key := strings.TrimPrefix(r.URL.Path, "/folder/")
|
|
||||||
|
|
||||||
switch r.Method {
|
|
||||||
case "PROPFIND":
|
|
||||||
s.mu.Lock()
|
|
||||||
var b strings.Builder
|
|
||||||
for name := range s.objects {
|
|
||||||
fmt.Fprintf(&b, "<d:href>/folder/%s</d:href>\n", name)
|
|
||||||
}
|
|
||||||
s.mu.Unlock()
|
|
||||||
w.Write([]byte(b.String()))
|
|
||||||
case http.MethodPut:
|
|
||||||
body, _ := io.ReadAll(r.Body)
|
|
||||||
s.mu.Lock()
|
|
||||||
s.objects[key] = body
|
|
||||||
s.mu.Unlock()
|
|
||||||
w.WriteHeader(http.StatusCreated)
|
|
||||||
case http.MethodGet:
|
|
||||||
s.mu.Lock()
|
|
||||||
data, ok := s.objects[key]
|
|
||||||
s.mu.Unlock()
|
|
||||||
if !ok {
|
|
||||||
w.WriteHeader(http.StatusNotFound)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
w.Write(data)
|
|
||||||
case http.MethodDelete:
|
|
||||||
s.mu.Lock()
|
|
||||||
delete(s.objects, key)
|
|
||||||
s.mu.Unlock()
|
|
||||||
w.WriteHeader(http.StatusNoContent)
|
|
||||||
default:
|
|
||||||
w.WriteHeader(http.StatusMethodNotAllowed)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *fakeStore) count() int {
|
|
||||||
s.mu.Lock()
|
|
||||||
defer s.mu.Unlock()
|
|
||||||
return len(s.objects)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *fakeStore) seenAuth() string {
|
|
||||||
s.mu.Lock()
|
|
||||||
defer s.mu.Unlock()
|
|
||||||
return s.sawAuth
|
|
||||||
}
|
|
||||||
|
|
||||||
func templateSettings(store *fakeStore, auth map[string]interface{}, secrets []string) *internet.MemoryStreamConfig {
|
|
||||||
base := store.server.URL
|
|
||||||
tmpl := map[string]interface{}{
|
|
||||||
"flatten": true,
|
|
||||||
"auth": auth,
|
|
||||||
"put": map[string]interface{}{"method": "PUT", "url": base + "/folder/{name}"},
|
|
||||||
"get": map[string]interface{}{"method": "GET", "url": base + "/folder/{name}"},
|
|
||||||
"delete": map[string]interface{}{"method": "DELETE", "url": base + "/folder/{name}"},
|
|
||||||
"list": map[string]interface{}{
|
|
||||||
"method": "PROPFIND",
|
|
||||||
"url": base + "/folder/",
|
|
||||||
"namesRegex": `<d:href>/folder/([^<]+)</d:href>`,
|
|
||||||
},
|
|
||||||
"retry": map[string]interface{}{"status": []int{429, 500, 502, 503}},
|
|
||||||
}
|
|
||||||
raw, _ := json.Marshal(tmpl)
|
|
||||||
return &internet.MemoryStreamConfig{
|
|
||||||
ProtocolName: protocolName,
|
|
||||||
ProtocolSettings: &Config{
|
|
||||||
RemoteFolder: "folder",
|
|
||||||
Service: "template",
|
|
||||||
Secrets: secrets,
|
|
||||||
Template: string(raw),
|
|
||||||
FlushIntervalMs: 5,
|
|
||||||
PollIntervalMs: 5,
|
|
||||||
MaxPollIntervalMs: 20,
|
|
||||||
SessionTtlSeconds: 5,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTemplateBackend(t *testing.T, store *fakeStore, auth map[string]interface{}, secrets []string) *templateStorage {
|
|
||||||
t.Helper()
|
|
||||||
settings := templateSettings(store, auth, secrets)
|
|
||||||
storage, err := newTemplateStorage(settings, settings.ProtocolSettings.(*Config))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newTemplateStorage: %v", err)
|
|
||||||
}
|
|
||||||
return storage
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTemplateRoundTrip(t *testing.T) {
|
|
||||||
store := newFakeStore(t)
|
|
||||||
storage := newTemplateBackend(t, store, map[string]interface{}{"type": "none"}, nil)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
if err := storage.Put(ctx, "streams/abc/c2s/000000000.seg", []byte("hello")); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
data, err := storage.Get(ctx, "streams/abc/c2s/000000000.seg")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Get: %v", err)
|
|
||||||
}
|
|
||||||
if string(data) != "hello" {
|
|
||||||
t.Fatalf("Get returned %q, want hello", data)
|
|
||||||
}
|
|
||||||
|
|
||||||
entries, err := storage.List(ctx, "streams/abc/c2s")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List: %v", err)
|
|
||||||
}
|
|
||||||
if len(entries) != 1 || entries[0].Name != "000000000.seg" {
|
|
||||||
t.Fatalf("List returned %v, want one segment", entries)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := storage.Get(ctx, "streams/abc/c2s/000000009.seg"); err != errNotFound {
|
|
||||||
t.Fatalf("Get of a missing object returned %v, want errNotFound", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := storage.Delete(ctx, "streams/abc/c2s/000000000.seg"); err != nil {
|
|
||||||
t.Fatalf("Delete: %v", err)
|
|
||||||
}
|
|
||||||
if store.count() != 0 {
|
|
||||||
t.Fatalf("store still holds %d objects", store.count())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTemplateListReturnsDirectChildren(t *testing.T) {
|
|
||||||
store := newFakeStore(t)
|
|
||||||
storage := newTemplateBackend(t, store, map[string]interface{}{"type": "none"}, nil)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
for _, name := range []string{
|
|
||||||
"streams/one/c2s/000000000.seg",
|
|
||||||
"streams/one/s2c/000000000.seg",
|
|
||||||
"streams/two/c2s/000000000.seg",
|
|
||||||
} {
|
|
||||||
if err := storage.Put(ctx, name, []byte("x")); err != nil {
|
|
||||||
t.Fatalf("Put %s: %v", name, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
entries, err := storage.List(ctx, "streams")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List: %v", err)
|
|
||||||
}
|
|
||||||
if len(entries) != 2 {
|
|
||||||
t.Fatalf("List returned %v, want the two session ids", entries)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTemplateBasicAuth(t *testing.T) {
|
|
||||||
store := newFakeStore(t)
|
|
||||||
store.needAuth = "Basic dXNlcjpwYXNz"
|
|
||||||
auth := map[string]interface{}{"type": "basic", "username": "{secret0}", "password": "{secret1}"}
|
|
||||||
storage := newTemplateBackend(t, store, auth, []string{"user", "pass"})
|
|
||||||
|
|
||||||
if err := storage.Put(context.Background(), "sessions/a", []byte("x")); err != nil {
|
|
||||||
t.Fatalf("Put with basic auth: %v", err)
|
|
||||||
}
|
|
||||||
if store.seenAuth() != "Basic dXNlcjpwYXNz" {
|
|
||||||
t.Fatalf("server saw auth %q", store.seenAuth())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTemplateOAuth(t *testing.T) {
|
|
||||||
store := newFakeStore(t)
|
|
||||||
store.needAuth = "Bearer tok-fake"
|
|
||||||
auth := map[string]interface{}{
|
|
||||||
"type": "oauth2",
|
|
||||||
"tokenUrl": store.server.URL + "/token",
|
|
||||||
"form": map[string]interface{}{"grant_type": "refresh_token", "refresh_token": "{secret0}"},
|
|
||||||
"header": map[string]interface{}{"Authorization": "Bearer {token}"},
|
|
||||||
}
|
|
||||||
storage := newTemplateBackend(t, store, auth, []string{"refresh"})
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
for i := 0; i < 4; i++ {
|
|
||||||
if err := storage.Put(ctx, fmt.Sprintf("sessions/s%d", i), nil); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if store.tokens != 1 {
|
|
||||||
t.Fatalf("token endpoint was hit %d times, want 1", store.tokens)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTemplateTransport(t *testing.T) {
|
|
||||||
store := newFakeStore(t)
|
|
||||||
settings := templateSettings(store, map[string]interface{}{"type": "none"}, nil)
|
|
||||||
|
|
||||||
client, server, cleanup := pairWith(t, settings)
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
if _, err := client.Write([]byte("ping")); err != nil {
|
|
||||||
t.Fatalf("client write: %v", err)
|
|
||||||
}
|
|
||||||
expectRead(t, server, "ping")
|
|
||||||
|
|
||||||
if _, err := server.Write([]byte("pong")); err != nil {
|
|
||||||
t.Fatalf("server write: %v", err)
|
|
||||||
}
|
|
||||||
expectRead(t, client, "pong")
|
|
||||||
|
|
||||||
payload := make([]byte, 300000)
|
|
||||||
for i := range payload {
|
|
||||||
payload[i] = byte(i % 251)
|
|
||||||
}
|
|
||||||
go func() { client.Write(payload) }()
|
|
||||||
|
|
||||||
if err := server.SetReadDeadline(time.Now().Add(30 * time.Second)); err != nil {
|
|
||||||
t.Fatalf("SetReadDeadline: %v", err)
|
|
||||||
}
|
|
||||||
got := make([]byte, len(payload))
|
|
||||||
if _, err := io.ReadFull(server, got); err != nil {
|
|
||||||
t.Fatalf("ReadFull: %v", err)
|
|
||||||
}
|
|
||||||
for i := range got {
|
|
||||||
if got[i] != payload[i] {
|
|
||||||
t.Fatalf("payload mismatch at byte %d", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTemplateConcurrency(t *testing.T) {
|
|
||||||
store := newFakeStore(t)
|
|
||||||
settings := templateSettings(store, map[string]interface{}{"type": "none"}, nil)
|
|
||||||
|
|
||||||
var tmpl map[string]interface{}
|
|
||||||
cfg := settings.ProtocolSettings.(*Config)
|
|
||||||
if err := json.Unmarshal([]byte(cfg.Template), &tmpl); err != nil {
|
|
||||||
t.Fatalf("unmarshal: %v", err)
|
|
||||||
}
|
|
||||||
tmpl["concurrency"] = 4
|
|
||||||
raw, _ := json.Marshal(tmpl)
|
|
||||||
cfg.Template = string(raw)
|
|
||||||
|
|
||||||
storage, err := newTemplateStorage(settings, cfg)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newTemplateStorage: %v", err)
|
|
||||||
}
|
|
||||||
if cap(storage.inflight) != 4 {
|
|
||||||
t.Fatalf("inflight cap is %d, want 4 from the template", cap(storage.inflight))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTemplateConcurrencyDefault(t *testing.T) {
|
|
||||||
store := newFakeStore(t)
|
|
||||||
storage := newTemplateBackend(t, store, map[string]interface{}{"type": "none"}, nil)
|
|
||||||
if cap(storage.inflight) != driveMaxInflight {
|
|
||||||
t.Fatalf("default inflight cap is %d, want %d", cap(storage.inflight), driveMaxInflight)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,445 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
maxCoalescedTicks = 8
|
|
||||||
|
|
||||||
segSuffix = ".seg"
|
|
||||||
endSuffix = ".end"
|
|
||||||
errSuffix = ".err"
|
|
||||||
)
|
|
||||||
|
|
||||||
func objectName(prefix string, seq int64, suffix string) string {
|
|
||||||
return fmt.Sprintf("%s/%09d%s", prefix, seq, suffix)
|
|
||||||
}
|
|
||||||
|
|
||||||
func parseEntry(name string) (int64, bool) {
|
|
||||||
dot := strings.LastIndexByte(name, '.')
|
|
||||||
if dot < 0 {
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
switch name[dot:] {
|
|
||||||
case segSuffix, endSuffix, errSuffix:
|
|
||||||
default:
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
seq, err := strconv.ParseInt(name[:dot], 10, 64)
|
|
||||||
if err != nil || seq < 0 {
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
return seq, true
|
|
||||||
}
|
|
||||||
|
|
||||||
type walWriter struct {
|
|
||||||
ctx context.Context
|
|
||||||
storage Storage
|
|
||||||
prefix string
|
|
||||||
params
|
|
||||||
sem chan struct{}
|
|
||||||
wg sync.WaitGroup
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
buf []byte
|
|
||||||
seq int64
|
|
||||||
lastSize int
|
|
||||||
held int
|
|
||||||
closed bool
|
|
||||||
err error
|
|
||||||
}
|
|
||||||
|
|
||||||
func newWALWriter(ctx context.Context, storage Storage, prefix string, p params) *walWriter {
|
|
||||||
w := &walWriter{
|
|
||||||
ctx: ctx,
|
|
||||||
storage: storage,
|
|
||||||
prefix: prefix,
|
|
||||||
params: p,
|
|
||||||
sem: make(chan struct{}, p.concurrency),
|
|
||||||
}
|
|
||||||
go w.flushLoop()
|
|
||||||
return w
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *walWriter) Write(p []byte) (int, error) {
|
|
||||||
w.mu.Lock()
|
|
||||||
defer w.mu.Unlock()
|
|
||||||
|
|
||||||
if w.err != nil {
|
|
||||||
return 0, w.err
|
|
||||||
}
|
|
||||||
if w.closed {
|
|
||||||
return 0, io.ErrClosedPipe
|
|
||||||
}
|
|
||||||
|
|
||||||
w.buf = append(w.buf, p...)
|
|
||||||
for len(w.buf) >= w.segmentBytes {
|
|
||||||
if err := w.flushLocked(); err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return len(p), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *walWriter) flushLoop() {
|
|
||||||
ticker := time.NewTicker(w.flushInterval)
|
|
||||||
defer ticker.Stop()
|
|
||||||
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-w.ctx.Done():
|
|
||||||
return
|
|
||||||
case <-ticker.C:
|
|
||||||
w.mu.Lock()
|
|
||||||
if !w.closed && w.err == nil && len(w.buf) > 0 && w.readyToFlush() {
|
|
||||||
w.flushLocked()
|
|
||||||
}
|
|
||||||
done := w.closed || w.err != nil
|
|
||||||
w.mu.Unlock()
|
|
||||||
if done {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *walWriter) readyToFlush() bool {
|
|
||||||
grew := len(w.buf) > w.lastSize
|
|
||||||
w.lastSize = len(w.buf)
|
|
||||||
|
|
||||||
if grew && w.held < maxCoalescedTicks {
|
|
||||||
w.held++
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
w.held = 0
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *walWriter) flushLocked() error {
|
|
||||||
if len(w.buf) == 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
if w.err != nil {
|
|
||||||
return w.err
|
|
||||||
}
|
|
||||||
|
|
||||||
n := len(w.buf)
|
|
||||||
if n > w.segmentBytes {
|
|
||||||
n = w.segmentBytes
|
|
||||||
}
|
|
||||||
|
|
||||||
chunk := make([]byte, n)
|
|
||||||
copy(chunk, w.buf[:n])
|
|
||||||
seq := w.seq
|
|
||||||
w.seq++
|
|
||||||
|
|
||||||
if n == len(w.buf) {
|
|
||||||
w.buf = w.buf[:0]
|
|
||||||
} else {
|
|
||||||
w.buf = append(w.buf[:0], w.buf[n:]...)
|
|
||||||
}
|
|
||||||
w.lastSize = len(w.buf)
|
|
||||||
w.held = 0
|
|
||||||
|
|
||||||
w.wg.Add(1)
|
|
||||||
go w.upload(seq, chunk)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *walWriter) upload(seq int64, chunk []byte) {
|
|
||||||
defer w.wg.Done()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case w.sem <- struct{}{}:
|
|
||||||
case <-w.ctx.Done():
|
|
||||||
return
|
|
||||||
}
|
|
||||||
defer func() { <-w.sem }()
|
|
||||||
|
|
||||||
if err := w.storage.Put(w.ctx, objectName(w.prefix, seq, segSuffix), chunk); err != nil {
|
|
||||||
w.mu.Lock()
|
|
||||||
if w.err == nil {
|
|
||||||
w.err = errors.New("failed to store segment").Base(err)
|
|
||||||
}
|
|
||||||
w.mu.Unlock()
|
|
||||||
w.storage.Put(w.ctx, objectName(w.prefix, seq, errSuffix), nil)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *walWriter) Close() error {
|
|
||||||
w.mu.Lock()
|
|
||||||
if w.closed {
|
|
||||||
w.mu.Unlock()
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
w.closed = true
|
|
||||||
for len(w.buf) > 0 && w.err == nil {
|
|
||||||
w.flushLocked()
|
|
||||||
}
|
|
||||||
w.mu.Unlock()
|
|
||||||
|
|
||||||
w.wg.Wait()
|
|
||||||
|
|
||||||
w.mu.Lock()
|
|
||||||
err, seq := w.err, w.seq
|
|
||||||
w.mu.Unlock()
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return w.storage.Put(w.ctx, objectName(w.prefix, seq, endSuffix), nil)
|
|
||||||
}
|
|
||||||
|
|
||||||
type walReader struct {
|
|
||||||
ctx context.Context
|
|
||||||
storage Storage
|
|
||||||
prefix string
|
|
||||||
params
|
|
||||||
seq int64
|
|
||||||
|
|
||||||
ch chan []byte
|
|
||||||
discards chan string
|
|
||||||
wake chan struct{}
|
|
||||||
holeSince time.Time
|
|
||||||
|
|
||||||
errMu sync.Mutex
|
|
||||||
err error
|
|
||||||
}
|
|
||||||
|
|
||||||
func newWALReader(ctx context.Context, storage Storage, prefix string, p params) *walReader {
|
|
||||||
r := &walReader{
|
|
||||||
ctx: ctx,
|
|
||||||
storage: storage,
|
|
||||||
prefix: prefix,
|
|
||||||
params: p,
|
|
||||||
ch: make(chan []byte, p.concurrency),
|
|
||||||
discards: make(chan string, 4*p.concurrency),
|
|
||||||
wake: make(chan struct{}, 1),
|
|
||||||
}
|
|
||||||
go r.run()
|
|
||||||
go r.discardLoop()
|
|
||||||
return r
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *walReader) Wake() {
|
|
||||||
select {
|
|
||||||
case r.wake <- struct{}{}:
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *walReader) run() {
|
|
||||||
defer close(r.ch)
|
|
||||||
|
|
||||||
delay := r.minPollInterval
|
|
||||||
active := time.Now()
|
|
||||||
for {
|
|
||||||
polled := time.Now()
|
|
||||||
advanced, eof, err := r.poll()
|
|
||||||
if err != nil {
|
|
||||||
r.setErr(err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if eof {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
switch {
|
|
||||||
case advanced:
|
|
||||||
active = time.Now()
|
|
||||||
delay = r.minPollInterval
|
|
||||||
case time.Since(active) < r.eagerWindow:
|
|
||||||
delay = r.minPollInterval
|
|
||||||
default:
|
|
||||||
delay *= 2
|
|
||||||
if delay > r.maxPollInterval {
|
|
||||||
delay = r.maxPollInterval
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
timer := time.NewTimer(delay)
|
|
||||||
select {
|
|
||||||
case <-r.ctx.Done():
|
|
||||||
timer.Stop()
|
|
||||||
return
|
|
||||||
case <-timer.C:
|
|
||||||
case <-r.wake:
|
|
||||||
timer.Stop()
|
|
||||||
active = time.Now()
|
|
||||||
delay = r.minPollInterval
|
|
||||||
if rest := r.minPollInterval - time.Since(polled); rest > 0 {
|
|
||||||
select {
|
|
||||||
case <-r.ctx.Done():
|
|
||||||
return
|
|
||||||
case <-time.After(rest):
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *walReader) poll() (advanced, eof bool, err error) {
|
|
||||||
listed, err := r.storage.List(r.ctx, r.prefix)
|
|
||||||
if err != nil {
|
|
||||||
return false, false, err
|
|
||||||
}
|
|
||||||
if len(listed) == 0 {
|
|
||||||
return false, false, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
pending := make(map[int64]Entry, len(listed))
|
|
||||||
ahead := false
|
|
||||||
for _, entry := range listed {
|
|
||||||
seq, ok := parseEntry(entry.Name)
|
|
||||||
if !ok {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
pending[seq] = entry
|
|
||||||
if seq > r.seq {
|
|
||||||
ahead = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, ok := pending[r.seq]; !ok && ahead {
|
|
||||||
if r.holeSince.IsZero() {
|
|
||||||
r.holeSince = time.Now()
|
|
||||||
} else if time.Since(r.holeSince) >= r.holeTimeout {
|
|
||||||
return false, false, errors.New("segment ", r.seq,
|
|
||||||
" never arrived while later ones did, the peer lost it")
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
r.holeSince = time.Time{}
|
|
||||||
}
|
|
||||||
|
|
||||||
for {
|
|
||||||
if entry, ok := pending[r.seq]; ok && strings.HasSuffix(entry.Name, errSuffix) {
|
|
||||||
r.discard(r.prefix + "/" + entry.Name)
|
|
||||||
return advanced, false, errors.New("the peer could not store segment ", r.seq)
|
|
||||||
}
|
|
||||||
|
|
||||||
batch, done := r.nextBatch(pending)
|
|
||||||
if done {
|
|
||||||
return advanced, true, nil
|
|
||||||
}
|
|
||||||
if len(batch) == 0 {
|
|
||||||
return advanced, false, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
chunks, err := r.fetch(batch)
|
|
||||||
if err != nil {
|
|
||||||
if err == errNotFound {
|
|
||||||
return advanced, false, nil
|
|
||||||
}
|
|
||||||
return advanced, false, err
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, chunk := range chunks {
|
|
||||||
select {
|
|
||||||
case r.ch <- chunk:
|
|
||||||
case <-r.ctx.Done():
|
|
||||||
return advanced, true, nil
|
|
||||||
}
|
|
||||||
r.seq++
|
|
||||||
advanced = true
|
|
||||||
r.discard(r.prefix + "/" + batch[i].Name)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *walReader) nextBatch(pending map[int64]Entry) (batch []Entry, done bool) {
|
|
||||||
for i := 0; i < r.concurrency; i++ {
|
|
||||||
entry, ok := pending[r.seq+int64(i)]
|
|
||||||
if !ok {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
if !strings.HasSuffix(entry.Name, segSuffix) {
|
|
||||||
if i == 0 && strings.HasSuffix(entry.Name, endSuffix) {
|
|
||||||
r.discard(r.prefix + "/" + entry.Name)
|
|
||||||
return nil, true
|
|
||||||
}
|
|
||||||
break
|
|
||||||
}
|
|
||||||
batch = append(batch, entry)
|
|
||||||
}
|
|
||||||
return batch, false
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *walReader) fetch(batch []Entry) ([][]byte, error) {
|
|
||||||
chunks := make([][]byte, len(batch))
|
|
||||||
failures := make([]error, len(batch))
|
|
||||||
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
for i, entry := range batch {
|
|
||||||
if entry.Inline != nil {
|
|
||||||
chunks[i] = entry.Inline
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
wg.Add(1)
|
|
||||||
go func(i int, name string) {
|
|
||||||
defer wg.Done()
|
|
||||||
chunks[i], failures[i] = r.storage.Get(r.ctx, r.prefix+"/"+name)
|
|
||||||
}(i, entry.Name)
|
|
||||||
}
|
|
||||||
wg.Wait()
|
|
||||||
|
|
||||||
for i := range batch {
|
|
||||||
if failures[i] != nil {
|
|
||||||
return nil, failures[i]
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return chunks, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *walReader) discard(name string) {
|
|
||||||
select {
|
|
||||||
case r.discards <- name:
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *walReader) discardLoop() {
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
defer wg.Wait()
|
|
||||||
|
|
||||||
sem := make(chan struct{}, r.concurrency)
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-r.ctx.Done():
|
|
||||||
return
|
|
||||||
case name := <-r.discards:
|
|
||||||
sem <- struct{}{}
|
|
||||||
wg.Add(1)
|
|
||||||
go func(name string) {
|
|
||||||
defer wg.Done()
|
|
||||||
defer func() { <-sem }()
|
|
||||||
r.storage.Delete(r.ctx, name)
|
|
||||||
}(name)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *walReader) setErr(err error) {
|
|
||||||
r.errMu.Lock()
|
|
||||||
defer r.errMu.Unlock()
|
|
||||||
if r.err == nil {
|
|
||||||
r.err = err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *walReader) Err() error {
|
|
||||||
r.errMu.Lock()
|
|
||||||
defer r.errMu.Unlock()
|
|
||||||
if r.err != nil {
|
|
||||||
return r.err
|
|
||||||
}
|
|
||||||
return io.EOF
|
|
||||||
}
|
|
||||||
@@ -1,308 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"crypto/rand"
|
|
||||||
"encoding/hex"
|
|
||||||
"fmt"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
"github.com/xtls/xray-core/transport/internet/stat"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
protocolName = "xdrive"
|
|
||||||
sessionsDir = "sessions"
|
|
||||||
streamsDir = "streams"
|
|
||||||
uplinkDir = "c2s"
|
|
||||||
downlinkDir = "s2c"
|
|
||||||
)
|
|
||||||
|
|
||||||
func init() {
|
|
||||||
common.Must(internet.RegisterProtocolConfigCreator(protocolName, func() interface{} {
|
|
||||||
return new(Config)
|
|
||||||
}))
|
|
||||||
common.Must(internet.RegisterTransportDialer(protocolName, Dial))
|
|
||||||
common.Must(internet.RegisterTransportListener(protocolName, Serve))
|
|
||||||
}
|
|
||||||
|
|
||||||
func newSessionID() (string, error) {
|
|
||||||
buf := make([]byte, 16)
|
|
||||||
if _, err := rand.Read(buf); err != nil {
|
|
||||||
return "", errors.New("failed to generate session id").Base(err)
|
|
||||||
}
|
|
||||||
return hex.EncodeToString(buf), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func announceName(session string, at time.Time) string {
|
|
||||||
return fmt.Sprintf("%s/%d-%s", sessionsDir, at.UnixNano(), session)
|
|
||||||
}
|
|
||||||
|
|
||||||
func parseAnnounce(entry string) (string, time.Time, bool) {
|
|
||||||
dash := strings.IndexByte(entry, '-')
|
|
||||||
if dash <= 0 || dash == len(entry)-1 {
|
|
||||||
return "", time.Time{}, false
|
|
||||||
}
|
|
||||||
nanos, err := strconv.ParseInt(entry[:dash], 10, 64)
|
|
||||||
if err != nil {
|
|
||||||
return "", time.Time{}, false
|
|
||||||
}
|
|
||||||
return entry[dash+1:], time.Unix(0, nanos), true
|
|
||||||
}
|
|
||||||
|
|
||||||
func sessionPrefix(session string) string {
|
|
||||||
return streamsDir + "/" + session
|
|
||||||
}
|
|
||||||
|
|
||||||
func uplinkPrefix(session string) string {
|
|
||||||
return sessionPrefix(session) + "/" + uplinkDir
|
|
||||||
}
|
|
||||||
|
|
||||||
func downlinkPrefix(session string) string {
|
|
||||||
return sessionPrefix(session) + "/" + downlinkDir
|
|
||||||
}
|
|
||||||
|
|
||||||
func streamConfig(streamSettings *internet.MemoryStreamConfig) (*Config, error) {
|
|
||||||
config, ok := streamSettings.ProtocolSettings.(*Config)
|
|
||||||
if !ok || config == nil {
|
|
||||||
return nil, errors.New("invalid protocol settings")
|
|
||||||
}
|
|
||||||
return config, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (stat.Connection, error) {
|
|
||||||
config, err := streamConfig(streamSettings)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
storage, err := newStorage(streamSettings)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
session, err := newSessionID()
|
|
||||||
if err != nil {
|
|
||||||
storage.Close()
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := storage.Put(ctx, announceName(session, time.Now()), nil); err != nil {
|
|
||||||
storage.Close()
|
|
||||||
return nil, errors.New("failed to announce session ", session).Base(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
errors.LogInfo(ctx, "opened session ", session)
|
|
||||||
|
|
||||||
return newConn(context.Background(), storage,
|
|
||||||
uplinkPrefix(session), downlinkPrefix(session), paramsFromConfig(config), func() {
|
|
||||||
storage.Close()
|
|
||||||
}), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type Listener struct {
|
|
||||||
ctx context.Context
|
|
||||||
cancel context.CancelFunc
|
|
||||||
storage Storage
|
|
||||||
addConn internet.ConnHandler
|
|
||||||
params
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
active map[string]bool
|
|
||||||
handled map[string]time.Time
|
|
||||||
idleSince map[string]time.Time
|
|
||||||
}
|
|
||||||
|
|
||||||
func Serve(ctx context.Context, address net.Address, port net.Port, streamSettings *internet.MemoryStreamConfig, addConn internet.ConnHandler) (internet.Listener, error) {
|
|
||||||
config, err := streamConfig(streamSettings)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
storage, err := newStorage(streamSettings)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
listenerCtx, cancel := context.WithCancel(context.Background())
|
|
||||||
listener := &Listener{
|
|
||||||
ctx: listenerCtx,
|
|
||||||
cancel: cancel,
|
|
||||||
storage: storage,
|
|
||||||
addConn: addConn,
|
|
||||||
params: paramsFromConfig(config),
|
|
||||||
active: make(map[string]bool),
|
|
||||||
handled: make(map[string]time.Time),
|
|
||||||
idleSince: make(map[string]time.Time),
|
|
||||||
}
|
|
||||||
|
|
||||||
go listener.acceptLoop(ctx)
|
|
||||||
go listener.collectLoop(ctx)
|
|
||||||
|
|
||||||
return listener, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *Listener) acceptLoop(logCtx context.Context) {
|
|
||||||
delay := l.minPollInterval
|
|
||||||
active := time.Now()
|
|
||||||
for {
|
|
||||||
accepted, err := l.acceptPending(logCtx)
|
|
||||||
if err != nil {
|
|
||||||
errors.LogWarningInner(logCtx, err, "failed to list sessions")
|
|
||||||
}
|
|
||||||
|
|
||||||
switch {
|
|
||||||
case accepted:
|
|
||||||
active = time.Now()
|
|
||||||
delay = l.minPollInterval
|
|
||||||
case time.Since(active) < l.eagerWindow:
|
|
||||||
delay = l.minPollInterval
|
|
||||||
default:
|
|
||||||
delay *= 2
|
|
||||||
if delay > l.maxPollInterval {
|
|
||||||
delay = l.maxPollInterval
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-l.ctx.Done():
|
|
||||||
return
|
|
||||||
case <-time.After(delay):
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *Listener) acceptPending(logCtx context.Context) (bool, error) {
|
|
||||||
sessions, err := l.storage.List(l.ctx, sessionsDir)
|
|
||||||
if err != nil {
|
|
||||||
return false, err
|
|
||||||
}
|
|
||||||
|
|
||||||
accepted := false
|
|
||||||
for _, listed := range sessions {
|
|
||||||
entry := listed.Name
|
|
||||||
full := sessionsDir + "/" + entry
|
|
||||||
|
|
||||||
session, at, ok := parseAnnounce(entry)
|
|
||||||
if !ok {
|
|
||||||
go l.drop(full)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if time.Since(at) > l.sessionTTL {
|
|
||||||
errors.LogInfo(logCtx, "dropping the stale announcement of session ", session)
|
|
||||||
go l.drop(full)
|
|
||||||
go l.drop(sessionPrefix(session))
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if !l.claim(session) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
go l.drop(full)
|
|
||||||
errors.LogInfo(logCtx, "accepted session ", session)
|
|
||||||
accepted = true
|
|
||||||
l.addConn(l.newSessionConn(session))
|
|
||||||
}
|
|
||||||
return accepted, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *Listener) drop(name string) {
|
|
||||||
l.storage.Delete(l.ctx, name)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *Listener) claim(session string) bool {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
if l.active[session] || !l.handled[session].IsZero() {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
l.active[session] = true
|
|
||||||
l.handled[session] = time.Now()
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *Listener) newSessionConn(session string) *Conn {
|
|
||||||
return newConn(l.ctx, l.storage,
|
|
||||||
downlinkPrefix(session), uplinkPrefix(session), l.params, func() {
|
|
||||||
l.mu.Lock()
|
|
||||||
delete(l.active, session)
|
|
||||||
l.mu.Unlock()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *Listener) collectLoop(logCtx context.Context) {
|
|
||||||
interval := l.sessionTTL / 2
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-l.ctx.Done():
|
|
||||||
return
|
|
||||||
case <-time.After(interval):
|
|
||||||
}
|
|
||||||
if err := l.collect(); err != nil {
|
|
||||||
errors.LogWarningInner(logCtx, err, "failed to collect abandoned sessions")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *Listener) collect() error {
|
|
||||||
sessions, err := l.storage.List(l.ctx, streamsDir)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
now := time.Now()
|
|
||||||
var expired []string
|
|
||||||
|
|
||||||
l.mu.Lock()
|
|
||||||
present := make(map[string]bool, len(sessions))
|
|
||||||
for _, listed := range sessions {
|
|
||||||
session := listed.Name
|
|
||||||
present[session] = true
|
|
||||||
if l.active[session] {
|
|
||||||
delete(l.idleSince, session)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
since, seen := l.idleSince[session]
|
|
||||||
if !seen {
|
|
||||||
l.idleSince[session] = now
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if now.Sub(since) >= l.sessionTTL {
|
|
||||||
expired = append(expired, session)
|
|
||||||
delete(l.idleSince, session)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
for session := range l.idleSince {
|
|
||||||
if !present[session] {
|
|
||||||
delete(l.idleSince, session)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
for session, at := range l.handled {
|
|
||||||
if !l.active[session] && now.Sub(at) >= l.sessionTTL {
|
|
||||||
delete(l.handled, session)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
l.mu.Unlock()
|
|
||||||
|
|
||||||
for _, session := range expired {
|
|
||||||
if err := l.storage.Delete(l.ctx, sessionPrefix(session)); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *Listener) Addr() net.Addr {
|
|
||||||
return placeholderAddr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *Listener) Close() error {
|
|
||||||
l.cancel()
|
|
||||||
return l.storage.Close()
|
|
||||||
}
|
|
||||||
@@ -1,583 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"crypto/rand"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
"strings"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
"github.com/xtls/xray-core/transport/internet/stat"
|
|
||||||
)
|
|
||||||
|
|
||||||
const testPatience = 30 * time.Second
|
|
||||||
|
|
||||||
func settings(folder string) *internet.MemoryStreamConfig {
|
|
||||||
return &internet.MemoryStreamConfig{
|
|
||||||
ProtocolName: protocolName,
|
|
||||||
ProtocolSettings: &Config{
|
|
||||||
RemoteFolder: folder,
|
|
||||||
Service: "local",
|
|
||||||
FlushIntervalMs: 5,
|
|
||||||
PollIntervalMs: 5,
|
|
||||||
MaxPollIntervalMs: 20,
|
|
||||||
SessionTtlSeconds: 5,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func pair(t *testing.T) (client, server stat.Connection, cleanup func()) {
|
|
||||||
t.Helper()
|
|
||||||
return pairWith(t, settings(t.TempDir()))
|
|
||||||
}
|
|
||||||
|
|
||||||
func pairWith(t *testing.T, streamSettings *internet.MemoryStreamConfig) (client, server stat.Connection, cleanup func()) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
accepted := make(chan stat.Connection, 1)
|
|
||||||
|
|
||||||
listener, err := Serve(context.Background(), net.LocalHostIP, net.Port(0), streamSettings, func(conn stat.Connection) {
|
|
||||||
accepted <- conn
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Serve: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
client, err = Dial(context.Background(), net.Destination{}, streamSettings)
|
|
||||||
if err != nil {
|
|
||||||
listener.Close()
|
|
||||||
t.Fatalf("Dial: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
select {
|
|
||||||
case server = <-accepted:
|
|
||||||
case <-time.After(testPatience):
|
|
||||||
client.Close()
|
|
||||||
listener.Close()
|
|
||||||
t.Fatal("listener did not accept the session")
|
|
||||||
}
|
|
||||||
|
|
||||||
return client, server, func() {
|
|
||||||
client.Close()
|
|
||||||
server.Close()
|
|
||||||
listener.Close()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func expectRead(t *testing.T, conn stat.Connection, want string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
if err := conn.SetReadDeadline(time.Now().Add(testPatience)); err != nil {
|
|
||||||
t.Fatalf("SetReadDeadline: %v", err)
|
|
||||||
}
|
|
||||||
buf := make([]byte, len(want))
|
|
||||||
if _, err := io.ReadFull(conn, buf); err != nil {
|
|
||||||
t.Fatalf("ReadFull: %v", err)
|
|
||||||
}
|
|
||||||
if string(buf) != want {
|
|
||||||
t.Fatalf("read %q, want %q", buf, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRoundTrip(t *testing.T) {
|
|
||||||
client, server, cleanup := pair(t)
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
if _, err := client.Write([]byte("ping")); err != nil {
|
|
||||||
t.Fatalf("client write: %v", err)
|
|
||||||
}
|
|
||||||
expectRead(t, server, "ping")
|
|
||||||
|
|
||||||
if _, err := server.Write([]byte("pong")); err != nil {
|
|
||||||
t.Fatalf("server write: %v", err)
|
|
||||||
}
|
|
||||||
expectRead(t, client, "pong")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestInterleaved(t *testing.T) {
|
|
||||||
client, server, cleanup := pair(t)
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
for i := 0; i < 20; i++ {
|
|
||||||
if _, err := client.Write([]byte("up")); err != nil {
|
|
||||||
t.Fatalf("client write %d: %v", i, err)
|
|
||||||
}
|
|
||||||
expectRead(t, server, "up")
|
|
||||||
|
|
||||||
if _, err := server.Write([]byte("down")); err != nil {
|
|
||||||
t.Fatalf("server write %d: %v", i, err)
|
|
||||||
}
|
|
||||||
expectRead(t, client, "down")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMultiSegmentTransfer(t *testing.T) {
|
|
||||||
client, server, cleanup := pair(t)
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
payload := make([]byte, 3*defaultSegmentBytes+1234)
|
|
||||||
if _, err := rand.Read(payload); err != nil {
|
|
||||||
t.Fatalf("rand: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
client.Write(payload)
|
|
||||||
}()
|
|
||||||
|
|
||||||
if err := server.SetReadDeadline(time.Now().Add(30 * time.Second)); err != nil {
|
|
||||||
t.Fatalf("SetReadDeadline: %v", err)
|
|
||||||
}
|
|
||||||
got := make([]byte, len(payload))
|
|
||||||
if _, err := io.ReadFull(server, got); err != nil {
|
|
||||||
t.Fatalf("ReadFull: %v", err)
|
|
||||||
}
|
|
||||||
if !bytes.Equal(got, payload) {
|
|
||||||
t.Fatal("payload mismatch")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCloseEOF(t *testing.T) {
|
|
||||||
client, server, cleanup := pair(t)
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
if _, err := client.Write([]byte("bye")); err != nil {
|
|
||||||
t.Fatalf("client write: %v", err)
|
|
||||||
}
|
|
||||||
if err := client.Close(); err != nil {
|
|
||||||
t.Fatalf("client close: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := server.SetReadDeadline(time.Now().Add(5 * time.Second)); err != nil {
|
|
||||||
t.Fatalf("SetReadDeadline: %v", err)
|
|
||||||
}
|
|
||||||
got, err := io.ReadAll(server)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("ReadAll: %v", err)
|
|
||||||
}
|
|
||||||
if string(got) != "bye" {
|
|
||||||
t.Fatalf("read %q, want %q", got, "bye")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestReadDeadline(t *testing.T) {
|
|
||||||
client, _, cleanup := pair(t)
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
if err := client.SetReadDeadline(time.Now().Add(100 * time.Millisecond)); err != nil {
|
|
||||||
t.Fatalf("SetReadDeadline: %v", err)
|
|
||||||
}
|
|
||||||
buf := make([]byte, 4)
|
|
||||||
if _, err := client.Read(buf); !os.IsTimeout(err) {
|
|
||||||
t.Fatalf("Read returned %v, want a timeout", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLocalNameEscape(t *testing.T) {
|
|
||||||
root := t.TempDir()
|
|
||||||
storage, err := newLocalStorage(root)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newLocalStorage: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := storage.Put(context.Background(), "../escaped", []byte("x")); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
if _, err := os.Stat(root + "/escaped"); err != nil {
|
|
||||||
t.Fatalf("name was not clamped inside the root: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLocalMissingObject(t *testing.T) {
|
|
||||||
storage, err := newLocalStorage(t.TempDir())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newLocalStorage: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := storage.Get(context.Background(), "nothing/here"); err != errNotFound {
|
|
||||||
t.Fatalf("Get returned %v, want errNotFound", err)
|
|
||||||
}
|
|
||||||
names, err := storage.List(context.Background(), "nothing")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List: %v", err)
|
|
||||||
}
|
|
||||||
if len(names) != 0 {
|
|
||||||
t.Fatalf("List returned %v, want none", names)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestResumeAfterIdle(t *testing.T) {
|
|
||||||
client, server, cleanup := pair(t)
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
if _, err := client.Write([]byte("first")); err != nil {
|
|
||||||
t.Fatalf("client write: %v", err)
|
|
||||||
}
|
|
||||||
expectRead(t, server, "first")
|
|
||||||
|
|
||||||
time.Sleep(200 * time.Millisecond)
|
|
||||||
|
|
||||||
if _, err := client.Write([]byte("second")); err != nil {
|
|
||||||
t.Fatalf("client write: %v", err)
|
|
||||||
}
|
|
||||||
expectRead(t, server, "second")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestParseEntry(t *testing.T) {
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
seq int64
|
|
||||||
ok bool
|
|
||||||
}{
|
|
||||||
{"000000000.seg", 0, true},
|
|
||||||
{"000000042.seg", 42, true},
|
|
||||||
{"000000007.end", 7, true},
|
|
||||||
{"000000001.tmp", 0, false},
|
|
||||||
{"notanumber.seg", 0, false},
|
|
||||||
{"000000001", 0, false},
|
|
||||||
}
|
|
||||||
for _, c := range cases {
|
|
||||||
seq, ok := parseEntry(c.name)
|
|
||||||
if ok != c.ok || (ok && seq != c.seq) {
|
|
||||||
t.Fatalf("parseEntry(%q) = %d, %v; want %d, %v", c.name, seq, ok, c.seq, c.ok)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestParamDefaults(t *testing.T) {
|
|
||||||
p := paramsFromConfig(&Config{})
|
|
||||||
if p.segmentBytes != defaultSegmentBytes || p.flushInterval != defaultFlushInterval {
|
|
||||||
t.Fatalf("defaults not applied: %+v", p)
|
|
||||||
}
|
|
||||||
|
|
||||||
p = paramsFromConfig(&Config{SegmentBytes: 1 << 30, PollIntervalMs: 400, MaxPollIntervalMs: 100})
|
|
||||||
if p.segmentBytes != maxSegmentBytes {
|
|
||||||
t.Fatalf("segmentBytes is %d, want %d", p.segmentBytes, maxSegmentBytes)
|
|
||||||
}
|
|
||||||
if p.maxPollInterval < p.minPollInterval {
|
|
||||||
t.Fatalf("maxPollInterval %v below minPollInterval %v", p.maxPollInterval, p.minPollInterval)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func waitFor(t *testing.T, what string, done func() bool) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
deadline := time.Now().Add(5 * time.Second)
|
|
||||||
for time.Now().Before(deadline) {
|
|
||||||
if done() {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
time.Sleep(5 * time.Millisecond)
|
|
||||||
}
|
|
||||||
t.Fatalf("timed out waiting for %s", what)
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTestListener(t *testing.T, folder string) *Listener {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
storage, err := newLocalStorage(folder)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newLocalStorage: %v", err)
|
|
||||||
}
|
|
||||||
return &Listener{
|
|
||||||
ctx: context.Background(),
|
|
||||||
storage: storage,
|
|
||||||
params: paramsFromConfig(&Config{SessionTtlSeconds: 1}),
|
|
||||||
active: make(map[string]bool),
|
|
||||||
handled: make(map[string]time.Time),
|
|
||||||
idleSince: make(map[string]time.Time),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCollectAbandoned(t *testing.T) {
|
|
||||||
folder := t.TempDir()
|
|
||||||
listener := newTestListener(t, folder)
|
|
||||||
|
|
||||||
if err := listener.storage.Put(context.Background(), uplinkPrefix("dead")+"/000000000.seg", []byte("x")); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := listener.collect(); err != nil {
|
|
||||||
t.Fatalf("collect: %v", err)
|
|
||||||
}
|
|
||||||
names, _ := listener.storage.List(context.Background(), streamsDir)
|
|
||||||
if len(names) != 1 {
|
|
||||||
t.Fatalf("first pass removed the session, got %v", names)
|
|
||||||
}
|
|
||||||
|
|
||||||
listener.idleSince["dead"] = time.Now().Add(-2 * time.Second)
|
|
||||||
if err := listener.collect(); err != nil {
|
|
||||||
t.Fatalf("collect: %v", err)
|
|
||||||
}
|
|
||||||
names, _ = listener.storage.List(context.Background(), streamsDir)
|
|
||||||
if len(names) != 0 {
|
|
||||||
t.Fatalf("abandoned session still there, got %v", names)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCollectKeepsActive(t *testing.T) {
|
|
||||||
folder := t.TempDir()
|
|
||||||
listener := newTestListener(t, folder)
|
|
||||||
listener.active["live"] = true
|
|
||||||
|
|
||||||
if err := listener.storage.Put(context.Background(), uplinkPrefix("live")+"/000000000.seg", []byte("x")); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
listener.idleSince["live"] = time.Now().Add(-2 * time.Second)
|
|
||||||
if err := listener.collect(); err != nil {
|
|
||||||
t.Fatalf("collect: %v", err)
|
|
||||||
}
|
|
||||||
names, _ := listener.storage.List(context.Background(), streamsDir)
|
|
||||||
if len(names) != 1 {
|
|
||||||
t.Fatalf("collected an active session, got %v", names)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestParseAnnounce(t *testing.T) {
|
|
||||||
session, at, ok := parseAnnounce("1757000000123456789-abc123")
|
|
||||||
if !ok || session != "abc123" || at.UnixNano() != 1757000000123456789 {
|
|
||||||
t.Fatalf("parseAnnounce returned %q, %v, %v", session, at.UnixNano(), ok)
|
|
||||||
}
|
|
||||||
for _, bad := range []string{"abc123", "-abc123", "1757000000-", "notanumber-abc"} {
|
|
||||||
if _, _, ok := parseAnnounce(bad); ok {
|
|
||||||
t.Fatalf("parseAnnounce accepted %q", bad)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestStaleAnnounce(t *testing.T) {
|
|
||||||
folder := t.TempDir()
|
|
||||||
listener := newTestListener(t, folder)
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
stale := announceName("ghost", time.Now().Add(-time.Hour))
|
|
||||||
if err := listener.storage.Put(ctx, stale, nil); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
if err := listener.storage.Put(ctx, uplinkPrefix("ghost")+"/000000000.seg", []byte("x")); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
accepted, err := listener.acceptPending(ctx)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("acceptPending: %v", err)
|
|
||||||
}
|
|
||||||
if accepted {
|
|
||||||
t.Fatal("accepted a stale announcement")
|
|
||||||
}
|
|
||||||
|
|
||||||
waitFor(t, "the stale announcement to be removed", func() bool {
|
|
||||||
names, _ := listener.storage.List(ctx, sessionsDir)
|
|
||||||
return len(names) == 0
|
|
||||||
})
|
|
||||||
waitFor(t, "the stale session data to be removed", func() bool {
|
|
||||||
names, _ := listener.storage.List(ctx, streamsDir)
|
|
||||||
return len(names) == 0
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFreshAnnounce(t *testing.T) {
|
|
||||||
folder := t.TempDir()
|
|
||||||
listener := newTestListener(t, folder)
|
|
||||||
listener.addConn = func(conn stat.Connection) { conn.Close() }
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
if err := listener.storage.Put(ctx, announceName("fresh", time.Now()), nil); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
accepted, err := listener.acceptPending(ctx)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("acceptPending: %v", err)
|
|
||||||
}
|
|
||||||
if !accepted {
|
|
||||||
t.Fatal("did not accept a fresh announcement")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAnnouncePrecision(t *testing.T) {
|
|
||||||
at := time.Unix(1757000000, int64(900*time.Millisecond))
|
|
||||||
entry := strings.TrimPrefix(announceName("abc123", at), sessionsDir+"/")
|
|
||||||
|
|
||||||
session, parsed, ok := parseAnnounce(entry)
|
|
||||||
if !ok || session != "abc123" {
|
|
||||||
t.Fatalf("parseAnnounce(%q) returned %q, %v", entry, session, ok)
|
|
||||||
}
|
|
||||||
if !parsed.Equal(at) {
|
|
||||||
t.Fatalf("timestamp came back as %v, want %v", parsed, at)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRecentAnnounceTTL(t *testing.T) {
|
|
||||||
folder := t.TempDir()
|
|
||||||
listener := newTestListener(t, folder)
|
|
||||||
listener.addConn = func(conn stat.Connection) { conn.Close() }
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
recent := time.Now().Add(-900 * time.Millisecond)
|
|
||||||
if err := listener.storage.Put(ctx, announceName("recent", recent), nil); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
accepted, err := listener.acceptPending(ctx)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("acceptPending: %v", err)
|
|
||||||
}
|
|
||||||
if !accepted {
|
|
||||||
t.Fatal("dropped an announcement younger than the TTL")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMissingSegment(t *testing.T) {
|
|
||||||
storage, err := newLocalStorage(t.TempDir())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newLocalStorage: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
p := paramsFromConfig(&Config{PollIntervalMs: 5, MaxPollIntervalMs: 20, HoleTimeoutMs: 200})
|
|
||||||
if err := storage.Put(ctx, objectName("hole", 1, segSuffix), []byte("second")); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
reader := newWALReader(ctx, storage, "hole", p)
|
|
||||||
select {
|
|
||||||
case _, ok := <-reader.ch:
|
|
||||||
if ok {
|
|
||||||
t.Fatal("delivered data past a missing segment")
|
|
||||||
}
|
|
||||||
case <-time.After(5 * time.Second):
|
|
||||||
t.Fatal("reader did not give up on a missing segment")
|
|
||||||
}
|
|
||||||
|
|
||||||
err = reader.Err()
|
|
||||||
if err == nil || err == io.EOF {
|
|
||||||
t.Fatalf("Err returned %v, want a failure", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestIdleStreamWaits(t *testing.T) {
|
|
||||||
storage, err := newLocalStorage(t.TempDir())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newLocalStorage: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
p := paramsFromConfig(&Config{PollIntervalMs: 5, MaxPollIntervalMs: 20, HoleTimeoutMs: 100})
|
|
||||||
reader := newWALReader(ctx, storage, "idle", p)
|
|
||||||
|
|
||||||
time.Sleep(400 * time.Millisecond)
|
|
||||||
if err := storage.Put(ctx, objectName("idle", 0, segSuffix), []byte("late")); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
select {
|
|
||||||
case data, ok := <-reader.ch:
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("the reader gave up on an idle stream: %v", reader.Err())
|
|
||||||
}
|
|
||||||
if string(data) != "late" {
|
|
||||||
t.Fatalf("read %q, want %q", data, "late")
|
|
||||||
}
|
|
||||||
case <-time.After(5 * time.Second):
|
|
||||||
t.Fatal("reader missed a late segment")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFailureMarker(t *testing.T) {
|
|
||||||
storage, err := newLocalStorage(t.TempDir())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newLocalStorage: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
p := paramsFromConfig(&Config{PollIntervalMs: 5, MaxPollIntervalMs: 20})
|
|
||||||
if err := storage.Put(ctx, objectName("broken", 0, segSuffix), []byte("first")); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
if err := storage.Put(ctx, objectName("broken", 1, errSuffix), nil); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
reader := newWALReader(ctx, storage, "broken", p)
|
|
||||||
|
|
||||||
select {
|
|
||||||
case data, ok := <-reader.ch:
|
|
||||||
if !ok || string(data) != "first" {
|
|
||||||
t.Fatalf("want the segment before the marker, got %q %v", data, ok)
|
|
||||||
}
|
|
||||||
case <-time.After(5 * time.Second):
|
|
||||||
t.Fatal("reader did not deliver the first segment")
|
|
||||||
}
|
|
||||||
|
|
||||||
select {
|
|
||||||
case _, ok := <-reader.ch:
|
|
||||||
if ok {
|
|
||||||
t.Fatal("delivered data past the failure marker")
|
|
||||||
}
|
|
||||||
case <-time.After(5 * time.Second):
|
|
||||||
t.Fatal("reader did not stop on the failure marker")
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := reader.Err(); err == nil || err == io.EOF {
|
|
||||||
t.Fatalf("Err returned %v, want a failure", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type inlineOnlyStorage struct {
|
|
||||||
Storage
|
|
||||||
gets int64
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *inlineOnlyStorage) List(ctx context.Context, prefix string) ([]Entry, error) {
|
|
||||||
return []Entry{{Name: "000000000" + segSuffix, Inline: []byte("carried by the listing")}}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *inlineOnlyStorage) Get(ctx context.Context, name string) ([]byte, error) {
|
|
||||||
atomic.AddInt64(&s.gets, 1)
|
|
||||||
return nil, errNotFound
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *inlineOnlyStorage) Delete(ctx context.Context, name string) error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestInlinePayload(t *testing.T) {
|
|
||||||
base, err := newLocalStorage(t.TempDir())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newLocalStorage: %v", err)
|
|
||||||
}
|
|
||||||
storage := &inlineOnlyStorage{Storage: base}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
reader := newWALReader(ctx, storage, "inline", paramsFromConfig(&Config{PollIntervalMs: 5}))
|
|
||||||
select {
|
|
||||||
case data, ok := <-reader.ch:
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("reader stopped: %v", reader.Err())
|
|
||||||
}
|
|
||||||
if string(data) != "carried by the listing" {
|
|
||||||
t.Fatalf("read %q, want the inline payload", data)
|
|
||||||
}
|
|
||||||
case <-time.After(5 * time.Second):
|
|
||||||
t.Fatal("reader did not deliver the inline payload")
|
|
||||||
}
|
|
||||||
|
|
||||||
if got := atomic.LoadInt64(&storage.gets); got != 0 {
|
|
||||||
t.Fatalf("called Get %d times for an inline payload", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,171 +0,0 @@
|
|||||||
package xdrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"crypto/rand"
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net/http"
|
|
||||||
"os"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
)
|
|
||||||
|
|
||||||
const yandexBase = "https://webdav.yandex.ru"
|
|
||||||
|
|
||||||
func envUint(name string, def uint32) uint32 {
|
|
||||||
if v := os.Getenv(name); v != "" {
|
|
||||||
var n uint32
|
|
||||||
fmt.Sscanf(v, "%d", &n)
|
|
||||||
if n > 0 {
|
|
||||||
return n
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return def
|
|
||||||
}
|
|
||||||
|
|
||||||
func liveYandexSettings(t *testing.T) (*internet.MemoryStreamConfig, string, func()) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
user := os.Getenv("XDRIVE_YANDEX_USER")
|
|
||||||
pass := os.Getenv("XDRIVE_YANDEX_PASS")
|
|
||||||
if user == "" || pass == "" {
|
|
||||||
t.Skip("set XDRIVE_YANDEX_USER and XDRIVE_YANDEX_PASS to run this test")
|
|
||||||
}
|
|
||||||
|
|
||||||
folder := fmt.Sprintf("xdrive-live-%d", time.Now().UnixNano())
|
|
||||||
dav := func(method, path string) int {
|
|
||||||
req, _ := http.NewRequest(method, yandexBase+path, nil)
|
|
||||||
req.SetBasicAuth(user, pass)
|
|
||||||
resp, err := http.DefaultClient.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("%s %s: %v", method, path, err)
|
|
||||||
}
|
|
||||||
resp.Body.Close()
|
|
||||||
return resp.StatusCode
|
|
||||||
}
|
|
||||||
if code := dav("MKCOL", "/"+folder); code != 201 && code != 405 {
|
|
||||||
t.Fatalf("MKCOL answered %d", code)
|
|
||||||
}
|
|
||||||
|
|
||||||
tmpl := map[string]interface{}{
|
|
||||||
"flatten": true,
|
|
||||||
"auth": map[string]interface{}{"type": "basic", "username": "{secret0}", "password": "{secret1}"},
|
|
||||||
"put": map[string]interface{}{"method": "PUT", "url": yandexBase + "/{folder}/{name}"},
|
|
||||||
"get": map[string]interface{}{"method": "GET", "url": yandexBase + "/{folder}/{name}"},
|
|
||||||
"delete": map[string]interface{}{"method": "DELETE", "url": yandexBase + "/{folder}/{name}"},
|
|
||||||
"list": map[string]interface{}{
|
|
||||||
"method": "PROPFIND", "url": yandexBase + "/{folder}/",
|
|
||||||
"headers": map[string]interface{}{"Depth": "1"}, "namesRegex": `<d:href>[^<]*/([^/<]+)</d:href>`,
|
|
||||||
},
|
|
||||||
"retry": map[string]interface{}{"status": []int{429, 500, 502, 503}},
|
|
||||||
}
|
|
||||||
raw, _ := json.Marshal(tmpl)
|
|
||||||
settings := &internet.MemoryStreamConfig{
|
|
||||||
ProtocolName: protocolName,
|
|
||||||
ProtocolSettings: &Config{
|
|
||||||
RemoteFolder: folder,
|
|
||||||
Service: "template",
|
|
||||||
Secrets: []string{user, pass},
|
|
||||||
Template: string(raw),
|
|
||||||
SegmentBytes: 262144,
|
|
||||||
FlushIntervalMs: 100,
|
|
||||||
PollIntervalMs: 300,
|
|
||||||
MaxPollIntervalMs: 1500,
|
|
||||||
SessionTtlSeconds: 120,
|
|
||||||
Concurrency: envUint("XDRIVE_LIVE_CONCURRENCY", 8),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
cleanup := func() { dav("DELETE", "/"+folder) }
|
|
||||||
return settings, folder, cleanup
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLiveYandexStorage(t *testing.T) {
|
|
||||||
settings, _, cleanup := liveYandexSettings(t)
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
storage, err := newTemplateStorage(settings, settings.ProtocolSettings.(*Config))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newTemplateStorage: %v", err)
|
|
||||||
}
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
name := "streams/live/c2s/000000000.seg"
|
|
||||||
payload := []byte("xdrive over real yandex webdav")
|
|
||||||
|
|
||||||
start := time.Now()
|
|
||||||
if err := storage.Put(ctx, name, payload); err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
t.Logf("Put took %v", time.Since(start))
|
|
||||||
|
|
||||||
start = time.Now()
|
|
||||||
names, err := storage.List(ctx, "streams/live/c2s")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List: %v", err)
|
|
||||||
}
|
|
||||||
t.Logf("List took %v", time.Since(start))
|
|
||||||
if len(names) != 1 || names[0].Name != "000000000.seg" {
|
|
||||||
t.Fatalf("List returned %v, want one segment", names)
|
|
||||||
}
|
|
||||||
|
|
||||||
start = time.Now()
|
|
||||||
got, err := storage.Get(ctx, name)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Get: %v", err)
|
|
||||||
}
|
|
||||||
t.Logf("Get took %v", time.Since(start))
|
|
||||||
if !bytes.Equal(got, payload) {
|
|
||||||
t.Fatalf("Get returned %q", got)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := storage.Get(ctx, "streams/live/c2s/000000009.seg"); err != errNotFound {
|
|
||||||
t.Fatalf("Get of a missing object returned %v, want errNotFound", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := storage.Delete(ctx, name); err != nil {
|
|
||||||
t.Fatalf("Delete: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLiveYandexTransport(t *testing.T) {
|
|
||||||
settings, _, cleanup := liveYandexSettings(t)
|
|
||||||
defer cleanup()
|
|
||||||
|
|
||||||
client, server, done := pairWith(t, settings)
|
|
||||||
defer done()
|
|
||||||
|
|
||||||
start := time.Now()
|
|
||||||
if _, err := client.Write([]byte("ping")); err != nil {
|
|
||||||
t.Fatalf("client write: %v", err)
|
|
||||||
}
|
|
||||||
expectRead(t, server, "ping")
|
|
||||||
t.Logf("client to server round took %v", time.Since(start))
|
|
||||||
|
|
||||||
size := 1000000
|
|
||||||
if raw := os.Getenv("XDRIVE_LIVE_BYTES"); raw != "" {
|
|
||||||
fmt.Sscanf(raw, "%d", &size)
|
|
||||||
}
|
|
||||||
payload := make([]byte, size)
|
|
||||||
rand.Read(payload)
|
|
||||||
|
|
||||||
start = time.Now()
|
|
||||||
go func() { client.Write(payload) }()
|
|
||||||
if err := server.SetReadDeadline(time.Now().Add(5 * time.Minute)); err != nil {
|
|
||||||
t.Fatalf("SetReadDeadline: %v", err)
|
|
||||||
}
|
|
||||||
got := make([]byte, len(payload))
|
|
||||||
if _, err := io.ReadFull(server, got); err != nil {
|
|
||||||
t.Fatalf("ReadFull: %v", err)
|
|
||||||
}
|
|
||||||
elapsed := time.Since(start)
|
|
||||||
if !bytes.Equal(got, payload) {
|
|
||||||
t.Fatal("payload mismatch")
|
|
||||||
}
|
|
||||||
t.Logf("%d bytes in %v -> %.1f KiB/s", len(payload), elapsed,
|
|
||||||
float64(len(payload))/1024/elapsed.Seconds())
|
|
||||||
}
|
|
||||||
Reference in New Issue
Block a user