diff --git a/transport/internet/finalmask/xdns/client.go b/transport/internet/finalmask/xdns/client.go index 73b265e5a..fb39b94a2 100644 --- a/transport/internet/finalmask/xdns/client.go +++ b/transport/internet/finalmask/xdns/client.go @@ -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 } diff --git a/transport/internet/finalmask/xdns/frag.go b/transport/internet/finalmask/xdns/frag.go index 444f405af..eb5254889 100644 --- a/transport/internet/finalmask/xdns/frag.go +++ b/transport/internet/finalmask/xdns/frag.go @@ -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) }