XDNS finalmask: Refine upload (#7095)

Fixes https://github.com/XTLS/Xray-core/pull/7090#issuecomment-6009651970
This commit is contained in:
LjhAUMEM
2026-10-08 04:58:49 +00:00
committed by GitHub
parent 9e55a6ed18
commit 37184948d4
2 changed files with 135 additions and 165 deletions
+121 -141
View File
@@ -47,7 +47,6 @@ type xdnsClient struct {
resolverIndex atomic.Uint32
readCh chan packet
sendCh chan []byte
poolCh chan struct{}
closeCh chan struct{}
wg sync.WaitGroup
@@ -101,7 +100,6 @@ func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) {
resolverSends: make([]atomic.Uint32, len(c.Resolvers)),
readCh: make(chan packet),
sendCh: make(chan []byte, 16),
poolCh: make(chan struct{}, pollLimit),
closeCh: make(chan struct{}),
}
@@ -118,6 +116,107 @@ func (c *xdnsClient) closed() bool {
}
}
func (c *xdnsClient) send(p []byte) {
domain := c.domains[mrand.Intn(len(c.domains))]
qtype := domain.types[mrand.Intn(len(domain.types))]
var buf [512]byte
var data [255]byte
send := func(p []byte) {
msg := dnsmessage.Message{
Header: dnsmessage.Header{
RecursionDesired: true,
},
Questions: []dnsmessage.Question{
{
Name: domain.Encode(p),
Type: dnsmessage.Type(qtype),
Class: dnsmessage.ClassINET,
},
},
}
if domain.edns0 > 0 {
msg.Additionals = []dnsmessage.Resource{
{
Header: dnsmessage.ResourceHeader{
Name: dnsmessage.MustNewName("."),
Type: dnsmessage.TypeOPT,
Class: dnsmessage.Class(domain.edns0),
TTL: 0,
},
Body: &dnsmessage.OPTResource{},
},
}
}
pack := common.Must2(msg.AppendPack(buf[:0]))
common.Must2(rand.Read(pack[:2]))
index := c.resolverIndex.Load()
cur := c.resolverSends[index].Add(1)
i := index
for {
i++
if i == uint32(len(c.resolvers)) {
i = 0
}
if i == index {
break
}
if cur > c.resolverSends[i].Load() {
break
}
}
c.resolverIndex.Store(i)
c.resolvers[index].Send(pack)
}
if len(p) == 0 {
copy(data[:], c.clientID[:])
data[0] |= TypeMap[qtype]
data[8] = 8
common.Must2(rand.Read(data[9:17]))
send(data[:17])
return
}
if len(p) <= domain.cap-12 {
copy(data[:], c.clientID[:])
data[0] |= TypeMap[qtype]
data[8] = 3
common.Must2(rand.Read(data[9:12]))
copy(data[12:], p)
send(data[:12+len(p)])
return
}
if len(p) <= 255*(domain.cap-15) {
copy(data[:], c.clientID[:])
data[0] |= TypeMap[qtype]
data[8] = 3 | 0xC0
common.Must2(rand.Read(data[9:12]))
fragID := byte(c.fragID.Add(1))
fragN := len(p) / (domain.cap - 15)
if len(p)%(domain.cap-15) > 0 {
fragN++
}
for i := range fragN {
data[12] = fragID
data[13] = byte(i)
data[14] = byte(fragN)
size := min(len(p), domain.cap-15)
copy(data[15:], p[:size])
send(data[:15+size])
p = p[size:]
}
return
}
errors.LogError(context.Background(), "err size ", len(p))
}
func (c *xdnsClient) read(buf []byte, addr net.Addr) bool {
msg := dnsmessage.Message{}
if err := msg.Unpack(buf); err != nil {
@@ -193,11 +292,10 @@ func (c *xdnsClient) run() {
}
c.wg.Add(1)
go c.send()
go c.poll()
c.wg.Wait()
close(c.readCh)
close(c.sendCh)
close(c.poolCh)
}
@@ -224,152 +322,36 @@ func (c *xdnsClient) recv(i int) {
}
}
func (c *xdnsClient) send() {
func (c *xdnsClient) poll() {
defer c.wg.Done()
var buf [512]byte
var data [255]byte
sendMsg := func(p []byte, domain *Domain, qtype uint16) {
msg := dnsmessage.Message{
Header: dnsmessage.Header{
RecursionDesired: true,
},
Questions: []dnsmessage.Question{
{
Name: domain.Encode(p),
Type: dnsmessage.Type(qtype),
Class: dnsmessage.ClassINET,
},
},
}
if domain.edns0 > 0 {
msg.Additionals = []dnsmessage.Resource{
{
Header: dnsmessage.ResourceHeader{
Name: dnsmessage.MustNewName("."),
Type: dnsmessage.TypeOPT,
Class: dnsmessage.Class(domain.edns0),
TTL: 0,
},
Body: &dnsmessage.OPTResource{},
},
}
}
pack := common.Must2(msg.AppendPack(buf[:0]))
common.Must2(rand.Read(pack[:2]))
index := c.resolverIndex.Load()
cur := c.resolverSends[index].Add(1)
i := index
for {
i++
if i == uint32(len(c.resolvers)) {
i = 0
}
if i == index {
break
}
if cur > c.resolverSends[i].Load() {
break
}
}
c.resolverIndex.Store(i)
c.resolvers[index].Send(pack)
select {
case <-c.closeCh:
case <-c.poolCh:
}
send := func(p []byte) {
domain := c.domains[mrand.Intn(len(c.domains))]
qtype := domain.types[mrand.Intn(len(domain.types))]
if len(p) == 0 {
copy(data[:], c.clientID[:])
data[0] |= TypeMap[qtype]
data[8] = 8
common.Must2(rand.Read(data[9:17]))
sendMsg(data[:17], domain, qtype)
return
}
if len(p) <= domain.cap-12 {
copy(data[:], c.clientID[:])
data[0] |= TypeMap[qtype]
data[8] = 3
common.Must2(rand.Read(data[9:12]))
copy(data[12:], p)
sendMsg(data[:12+len(p)], domain, qtype)
return
}
if len(p) <= 255*(domain.cap-15) {
copy(data[:], c.clientID[:])
data[0] |= TypeMap[qtype]
data[8] = 3 | 0xC0
common.Must2(rand.Read(data[9:12]))
fragID := byte(c.fragID.Add(1))
fragN := len(p) / (domain.cap - 15)
if len(p)%(domain.cap-15) > 0 {
fragN++
}
for i := range fragN {
data[12] = fragID
data[13] = byte(i)
data[14] = byte(fragN)
size := min(len(p), domain.cap-15)
copy(data[15:], p[:size])
sendMsg(data[:15+size], domain, qtype)
p = p[size:]
}
return
}
errors.LogError(context.Background(), "err size ", len(p))
}
ticker := time.NewTicker(initPollDelay)
defer ticker.Stop()
delay := initPollDelay
p := []byte(nil)
timeout := false
ticker := time.NewTicker(delay)
defer ticker.Stop()
for {
select {
case <-c.closeCh:
return
default:
select {
case <-c.closeCh:
return
case p = <-c.sendCh:
case <-c.poolCh:
case <-ticker.C:
timeout = true
}
}
if len(p) > 0 {
select {
case <-c.poolCh:
default:
}
}
send(p)
for range c.extraPoll {
send(nil)
}
if timeout {
case <-c.poolCh:
delay = initPollDelay
case <-ticker.C:
delay *= pollDelayMultiplier
if delay > maxPollDelay {
delay = maxPollDelay
}
timeout = false
} else {
delay = initPollDelay
}
if c.closed() {
return
}
ticker.Reset(delay)
c.send(nil)
for range c.extraPoll {
c.send(nil)
}
}
}
@@ -391,11 +373,9 @@ func (c *xdnsClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
errors.LogError(context.Background(), "err size ", len(p))
return 0, errors.New("err size")
}
b := make([]byte, len(p))
copy(b, p)
select {
case c.sendCh <- b:
default:
c.send(p)
for range c.extraPoll {
c.send(nil)
}
return len(p), nil
}
+14 -24
View File
@@ -6,10 +6,9 @@ import (
)
const (
fragTTL = 8 * time.Second
fragSize = 4096
fragClientIDSize = 16384
fragCount = 4096
fragTTL = 4 * time.Second
fragSize = 4096
fragCount = 4096
)
type FragKey struct {
@@ -26,17 +25,15 @@ type FragEntry struct {
}
type FragManager struct {
m map[FragKey]*FragEntry
sizem map[ClientID]int
ch chan struct{}
mu sync.Mutex
m map[FragKey]*FragEntry
ch chan struct{}
mu sync.Mutex
}
func NewFragManager() *FragManager {
m := &FragManager{
m: make(map[FragKey]*FragEntry),
sizem: make(map[ClientID]int),
ch: make(chan struct{}),
m: make(map[FragKey]*FragEntry),
ch: make(chan struct{}),
}
go m.gc()
return m
@@ -51,9 +48,8 @@ func (m *FragManager) closed() bool {
}
}
func (m *FragManager) removeEntey(k FragKey, e *FragEntry) {
m.sizem[k.clientID] -= e.size
delete(m.m, k)
func (m *FragManager) remove(key FragKey) {
delete(m.m, key)
}
func (m *FragManager) tryRemove() {
@@ -70,7 +66,7 @@ func (m *FragManager) tryRemove() {
first = false
}
}
m.removeEntey(key, entry)
m.remove(key)
}
func (m *FragManager) gc() {
@@ -84,7 +80,7 @@ func (m *FragManager) gc() {
m.mu.Lock()
for k, e := range m.m {
if now.After(e.deadline) {
m.removeEntey(k, e)
m.remove(k)
}
}
m.mu.Unlock()
@@ -109,7 +105,7 @@ func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []
if entry == nil {
m.tryRemove()
} else {
m.removeEntey(key, entry)
m.remove(key)
}
entry = &FragEntry{
data: make([][]byte, fragN),
@@ -131,11 +127,6 @@ func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []
if entry.size+len(data) > fragSize {
return 0
}
if entry.len < int(entry.total)-1 {
if m.sizem[key.clientID]+len(data) > fragClientIDSize {
return 0
}
}
cp := make([]byte, len(data))
copy(cp, data)
@@ -144,7 +135,6 @@ func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []
entry.size += len(data)
entry.len++
entry.deadline = now.Add(fragTTL)
m.sizem[key.clientID] += len(data)
if entry.len < int(entry.total) {
return 0
@@ -154,7 +144,7 @@ func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []
for i := range entry.data {
out = append(out, entry.data[i]...)
}
m.removeEntey(key, entry)
m.remove(key)
return len(out)
}