mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-09-30 04:55:42 +00:00
XDNS finalmask: Refactor and new parameters (#6718)
https://github.com/XTLS/Xray-core/pull/6718#issuecomment-5894987590 Fixes https://github.com/XTLS/Xray-core/issues/6692
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
package conf
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/x509"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
@@ -14,6 +15,7 @@ import (
|
||||
googleuuid "github.com/google/uuid"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"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/mkcp/aes128gcm"
|
||||
@@ -81,7 +83,7 @@ var (
|
||||
"noise": func() interface{} { return new(NoiseMask) },
|
||||
"salamander": func() interface{} { return new(Salamander) },
|
||||
"sudoku": func() interface{} { return new(Sudoku) },
|
||||
"xdns": func() interface{} { return new(Xdns) },
|
||||
"xdns": func() interface{} { return new(XDNS) },
|
||||
"xicmp": func() interface{} { return new(Xicmp) },
|
||||
"realm": func() interface{} { return new(Realm) },
|
||||
"udphop": func() interface{} { return new(UDPHop) },
|
||||
@@ -694,32 +696,88 @@ func (c *Sudoku) Build() (proto.Message, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
type Xdns struct {
|
||||
Domain json.RawMessage `json:"domain"`
|
||||
|
||||
Domains []string `json:"domains"`
|
||||
Resolvers []string `json:"resolvers"`
|
||||
type XDNSDomain struct {
|
||||
Name string `json:"name"`
|
||||
LenLimit int32 `json:"lenLimit"`
|
||||
LabelLimit int32 `json:"labelLimit"`
|
||||
Types []int32 `json:"types"`
|
||||
Edns0 int32 `json:"edns0"`
|
||||
}
|
||||
|
||||
func (c *Xdns) Build() (proto.Message, error) {
|
||||
if c.Domain != nil {
|
||||
return nil, errors.PrintRemovedFeatureError("domain", "domains(server) & resolvers(client)")
|
||||
}
|
||||
type XDNSResolverTCP struct {
|
||||
Addr string `json:"addr"`
|
||||
}
|
||||
|
||||
if len(c.Domains) == 0 && len(c.Resolvers) == 0 {
|
||||
return nil, errors.New("empty domains & empty resolvers")
|
||||
}
|
||||
func (c *XDNSResolverTCP) Build() (proto.Message, error) {
|
||||
return &xdns.TCPResolverProto{Addr: c.Addr}, nil
|
||||
}
|
||||
|
||||
for _, r := range c.Resolvers {
|
||||
if !strings.Contains(r, "+udp://") {
|
||||
return nil, errors.New("invalid resolver ", r)
|
||||
type XDNSResolverUDP struct {
|
||||
Addr string `json:"addr"`
|
||||
}
|
||||
|
||||
func (c *XDNSResolverUDP) Build() (proto.Message, error) {
|
||||
return &xdns.UDPResolverProto{Addr: c.Addr}, nil
|
||||
}
|
||||
|
||||
var xdnsLoader = NewJSONConfigLoader(ConfigCreatorCache{
|
||||
"tcp": func() interface{} { return new(XDNSResolverTCP) },
|
||||
"udp": func() interface{} { return new(XDNSResolverUDP) },
|
||||
}, "type", "settings")
|
||||
|
||||
type XDNSResolver struct {
|
||||
Type string `json:"type"`
|
||||
Settings json.RawMessage `json:"settings"`
|
||||
}
|
||||
|
||||
type XDNS struct {
|
||||
Domains []XDNSDomain `json:"domains"`
|
||||
Resolvers []XDNSResolver `json:"resolvers"`
|
||||
ExtraPoll int32 `json:"extraPoll"`
|
||||
}
|
||||
|
||||
func (c *XDNS) Build() (proto.Message, error) {
|
||||
var domains []*xdns.DomainProto
|
||||
var resolvers []*serial.TypedMessage
|
||||
for i := range c.Domains {
|
||||
if c.Domains[i].LenLimit == 0 {
|
||||
c.Domains[i].LenLimit = 255
|
||||
}
|
||||
if c.Domains[i].LabelLimit == 0 {
|
||||
c.Domains[i].LabelLimit = 63
|
||||
}
|
||||
types := make([]uint16, 0, len(c.Domains[i].Types))
|
||||
for j := range c.Domains[i].Types {
|
||||
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||
}
|
||||
domain, err := xdns.NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
errors.LogInfo(context.Background(), domain.Show())
|
||||
domains = append(domains, &xdns.DomainProto{
|
||||
Name: c.Domains[i].Name,
|
||||
LenLimit: c.Domains[i].LenLimit,
|
||||
LabelLimit: c.Domains[i].LabelLimit,
|
||||
Types: c.Domains[i].Types,
|
||||
Edns0: c.Domains[i].Edns0,
|
||||
})
|
||||
}
|
||||
|
||||
return &xdns.Config{
|
||||
Domains: c.Domains,
|
||||
Resolvers: c.Resolvers,
|
||||
}, nil
|
||||
for i := range c.Resolvers {
|
||||
config, err := xdnsLoader.LoadWithID(c.Resolvers[i].Settings, c.Resolvers[i].Type)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pm, err := config.(interface{ Build() (proto.Message, error) }).Build()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resolvers = append(resolvers, serial.ToTypedMessage(pm))
|
||||
}
|
||||
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
|
||||
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
|
||||
}
|
||||
return &xdns.Config{Domains: domains, Resolvers: resolvers, ExtraPoll: c.ExtraPoll}, nil
|
||||
}
|
||||
|
||||
type XMC struct {
|
||||
|
||||
@@ -223,13 +223,6 @@ func (c *udpHopConn) Close() error {
|
||||
}
|
||||
_ = c.cur.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case packet := <-c.readCh:
|
||||
if packet.p != nil {
|
||||
pool.Put(packet.p[:cap(packet.p)])
|
||||
}
|
||||
default:
|
||||
}
|
||||
close(c.readCh)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1,417 +1,441 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base32"
|
||||
"encoding/binary"
|
||||
go_errors "errors"
|
||||
"io"
|
||||
"net"
|
||||
"strconv"
|
||||
mrand "math/rand"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"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/finalmask"
|
||||
"golang.org/x/net/dns/dnsmessage"
|
||||
)
|
||||
|
||||
const (
|
||||
numPadding = 3
|
||||
numPaddingForPoll = 8
|
||||
initPollDelay = 500 * time.Millisecond
|
||||
maxPollDelay = 10 * time.Second
|
||||
pollDelayMultiplier = 2.0
|
||||
pollLimit = 16
|
||||
)
|
||||
|
||||
var base32Encoding = base32.StdEncoding.WithPadding(base32.NoPadding)
|
||||
var pool4K = sync.Pool{
|
||||
New: func() any {
|
||||
return make([]byte, 4096)
|
||||
},
|
||||
}
|
||||
|
||||
type packet struct {
|
||||
p []byte
|
||||
addr net.Addr
|
||||
}
|
||||
|
||||
type xdnsConnClient struct {
|
||||
net.PacketConn
|
||||
type xdnsClient struct {
|
||||
dialer *finalmask.Dialer
|
||||
|
||||
resolverAddrs []*net.UDPAddr
|
||||
resolverTypes []uint16
|
||||
resolverIdx uint32
|
||||
resolverSend map[string]*atomic.Uint32
|
||||
clientID ClientID
|
||||
fragID atomic.Uint32
|
||||
domains []*Domain
|
||||
extraPoll int32
|
||||
|
||||
clientID []byte
|
||||
domains []Name
|
||||
resolvers []Resolver
|
||||
resolverSends []atomic.Uint32
|
||||
resolverIndex atomic.Uint32
|
||||
|
||||
pollChan chan struct{}
|
||||
readQueue chan *packet
|
||||
writeQueue chan *packet
|
||||
|
||||
closed bool
|
||||
mutex sync.Mutex
|
||||
readCh chan packet
|
||||
sendCh chan []byte
|
||||
poolCh chan struct{}
|
||||
closeCh chan struct{}
|
||||
wg sync.WaitGroup
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
if len(c.Domains) == 0 {
|
||||
return nil, errors.New("empty domains")
|
||||
}
|
||||
if len(c.Resolvers) == 0 {
|
||||
return nil, errors.New("empty resolvers")
|
||||
}
|
||||
|
||||
var domains []Name
|
||||
var servers []string
|
||||
var resolverTypes []uint16
|
||||
for _, rs := range c.Resolvers {
|
||||
domain, server, resolverType, err := parseResolver(rs)
|
||||
if err != nil {
|
||||
return nil, errors.New("invalid resolvers").Base(err)
|
||||
}
|
||||
domains = append(domains, domain)
|
||||
servers = append(servers, server)
|
||||
resolverTypes = append(resolverTypes, resolverType)
|
||||
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
|
||||
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
|
||||
}
|
||||
|
||||
var resolverAddrs []*net.UDPAddr
|
||||
resolverSend := make(map[string]*atomic.Uint32)
|
||||
for _, rs := range servers {
|
||||
h, p, err := net.SplitHostPort(rs)
|
||||
domains := make([]*Domain, 0, len(c.Domains))
|
||||
for i := range c.Domains {
|
||||
types := make([]uint16, 0, len(c.Domains[i].Types))
|
||||
for j := range c.Domains[i].Types {
|
||||
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||
}
|
||||
domain, err := NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ip := net.ParseIP(h)
|
||||
if ip == nil {
|
||||
return nil, errors.New("invalid ip address")
|
||||
}
|
||||
port, err := strconv.Atoi(p)
|
||||
domains = append(domains, domain)
|
||||
}
|
||||
resolvers := make([]Resolver, 0, len(c.Resolvers))
|
||||
for i := range c.Resolvers {
|
||||
resolver, err := NewResolver(c.Resolvers[i], dialer)
|
||||
if err != nil {
|
||||
return nil, errors.New("invalid port").Base(err)
|
||||
return nil, err
|
||||
}
|
||||
addr := &net.UDPAddr{IP: ip, Port: port}
|
||||
resolverAddrs = append(resolverAddrs, addr)
|
||||
resolverSend[addr.String()] = &atomic.Uint32{}
|
||||
resolvers = append(resolvers, resolver)
|
||||
}
|
||||
client := &xdnsClient{
|
||||
dialer: dialer,
|
||||
|
||||
conn := &xdnsConnClient{
|
||||
PacketConn: raw,
|
||||
clientID: NewClientID(),
|
||||
domains: domains,
|
||||
extraPoll: c.ExtraPoll,
|
||||
|
||||
resolverAddrs: resolverAddrs,
|
||||
resolverTypes: resolverTypes,
|
||||
resolverIdx: 0,
|
||||
resolverSend: resolverSend,
|
||||
resolvers: resolvers,
|
||||
resolverSends: make([]atomic.Uint32, len(c.Resolvers)),
|
||||
|
||||
clientID: make([]byte, 8),
|
||||
domains: domains,
|
||||
|
||||
pollChan: make(chan struct{}, pollLimit),
|
||||
readQueue: make(chan *packet, 256),
|
||||
writeQueue: make(chan *packet, 256),
|
||||
readCh: make(chan packet),
|
||||
sendCh: make(chan []byte, 16),
|
||||
poolCh: make(chan struct{}, pollLimit),
|
||||
closeCh: make(chan struct{}),
|
||||
}
|
||||
|
||||
common.Must2(rand.Read(conn.clientID))
|
||||
|
||||
go conn.recvLoop()
|
||||
go conn.sendLoop()
|
||||
|
||||
return conn, nil
|
||||
go client.run()
|
||||
return client, nil
|
||||
}
|
||||
|
||||
func (c *xdnsConnClient) recvLoop() {
|
||||
var buf [finalmask.UDPSize]byte
|
||||
func (c *xdnsClient) closed() bool {
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
for {
|
||||
if c.closed {
|
||||
func (c *xdnsClient) read(buf []byte, addr net.Addr) bool {
|
||||
msg := dnsmessage.Message{}
|
||||
if err := msg.Unpack(buf); err != nil {
|
||||
return false
|
||||
}
|
||||
if !msg.Header.Response || msg.Header.Truncated || msg.Header.RCode != dnsmessage.RCodeSuccess || len(msg.Questions) != 1 {
|
||||
return false
|
||||
}
|
||||
|
||||
var domain *Domain
|
||||
for i := range c.domains {
|
||||
if c.domains[i].IsDomain(msg.Questions[0].Name) {
|
||||
domain = c.domains[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
if domain == nil || !domain.HasType(uint16(msg.Questions[0].Type)) {
|
||||
return false
|
||||
}
|
||||
|
||||
n, addr, err := c.PacketConn.ReadFrom(buf[:])
|
||||
edns0 := uint16(0)
|
||||
for i := range msg.Additionals {
|
||||
if msg.Additionals[i].Header.Type == dnsmessage.TypeOPT {
|
||||
edns0 = uint16(msg.Additionals[i].Header.Class)
|
||||
break
|
||||
}
|
||||
}
|
||||
errors.LogDebug(context.Background(), addr, " edns0 ", edns0, " buf ", len(buf), " ", msg.Questions[0].Type)
|
||||
|
||||
resp := NewResp(msg, domain, 0)
|
||||
|
||||
p := pool4K.Get().([]byte)
|
||||
n := resp.Decode(p)
|
||||
p = p[:n]
|
||||
|
||||
b := p
|
||||
var bs [][]byte
|
||||
for len(b) > 1 {
|
||||
last := b[0]&0xC0 == 0xC0
|
||||
length := int(b[0]&0x3F)<<8 | int(b[1])
|
||||
b = b[2:]
|
||||
if length > len(b) {
|
||||
bs = nil
|
||||
break
|
||||
}
|
||||
packet := make([]byte, length)
|
||||
copy(packet, b)
|
||||
bs = append(bs, packet)
|
||||
if last {
|
||||
break
|
||||
}
|
||||
b = b[length:]
|
||||
if len(b) < 2 {
|
||||
bs = nil
|
||||
}
|
||||
}
|
||||
pool4K.Put(p[:cap(p)])
|
||||
|
||||
for i := range bs {
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
return true
|
||||
case c.readCh <- packet{p: bs[i], addr: addr}:
|
||||
}
|
||||
}
|
||||
return len(bs) > 0
|
||||
}
|
||||
|
||||
func (c *xdnsClient) run() {
|
||||
for i := range len(c.resolvers) {
|
||||
c.wg.Add(1)
|
||||
go c.recv(i)
|
||||
}
|
||||
|
||||
c.wg.Add(1)
|
||||
go c.send()
|
||||
|
||||
c.wg.Wait()
|
||||
close(c.readCh)
|
||||
close(c.sendCh)
|
||||
close(c.poolCh)
|
||||
}
|
||||
|
||||
func (c *xdnsClient) recv(i int) {
|
||||
defer c.wg.Done()
|
||||
|
||||
var buf [4096]byte
|
||||
for {
|
||||
n, err := c.resolvers[i].Read(buf[:])
|
||||
if err != nil {
|
||||
if go_errors.Is(err, net.ErrClosed) {
|
||||
break
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
continue
|
||||
errors.LogErrorInner(context.Background(), err, "recv err ", i)
|
||||
return
|
||||
}
|
||||
|
||||
if addr == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
send := c.resolverSend[addr.String()]
|
||||
if send == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
resp, err := MessageFromWireFormat(buf[:n])
|
||||
if err != nil {
|
||||
errors.LogDebug(context.Background(), addr, " xdns from wireformat err ", err)
|
||||
continue
|
||||
}
|
||||
|
||||
payload := dnsResponsePayload(&resp, c.domains)
|
||||
|
||||
r := bytes.NewReader(payload)
|
||||
anyPacket := false
|
||||
for {
|
||||
p, err := nextPacket(r)
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
anyPacket = true
|
||||
|
||||
buf := make([]byte, len(p))
|
||||
copy(buf, p)
|
||||
if c.read(buf[:n], c.resolvers[i].Addr()) {
|
||||
c.resolverSends[i].Store(0)
|
||||
select {
|
||||
case c.readQueue <- &packet{
|
||||
p: buf,
|
||||
addr: addr,
|
||||
}:
|
||||
default:
|
||||
errors.LogDebug(context.Background(), addr, " mask read err queue full")
|
||||
}
|
||||
}
|
||||
|
||||
if anyPacket {
|
||||
send.Store(0)
|
||||
select {
|
||||
case c.pollChan <- struct{}{}:
|
||||
case c.poolCh <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
errors.LogDebug(context.Background(), "xdns closed")
|
||||
|
||||
close(c.pollChan)
|
||||
close(c.readQueue)
|
||||
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
c.closed = true
|
||||
close(c.writeQueue)
|
||||
}
|
||||
|
||||
func (c *xdnsConnClient) sendLoop() {
|
||||
pollDelay := initPollDelay
|
||||
pollTimer := time.NewTimer(pollDelay)
|
||||
for {
|
||||
var p *packet
|
||||
pollTimerExpired := false
|
||||
func (c *xdnsClient) send() {
|
||||
defer c.wg.Done()
|
||||
|
||||
select {
|
||||
case p = <-c.writeQueue:
|
||||
default:
|
||||
select {
|
||||
case p = <-c.writeQueue:
|
||||
case <-c.pollChan:
|
||||
case <-pollTimer.C:
|
||||
pollTimerExpired = true
|
||||
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]))
|
||||
|
||||
if p != nil {
|
||||
select {
|
||||
case <-c.pollChan:
|
||||
default:
|
||||
index := c.resolverIndex.Load()
|
||||
cur := c.resolverSends[index].Add(1)
|
||||
i := index
|
||||
for {
|
||||
i++
|
||||
if i == uint32(len(c.resolvers)) {
|
||||
i = 0
|
||||
}
|
||||
} else {
|
||||
encoded, _ := encode(nil, c.clientID, c.domains[c.resolverIdx], c.resolverTypes[c.resolverIdx])
|
||||
p = &packet{
|
||||
p: encoded,
|
||||
if i == index {
|
||||
break
|
||||
}
|
||||
if cur > c.resolverSends[i].Load() {
|
||||
break
|
||||
}
|
||||
}
|
||||
c.resolverIndex.Store(i)
|
||||
c.resolvers[index].Send(pack)
|
||||
}
|
||||
|
||||
if pollTimerExpired {
|
||||
pollDelay = time.Duration(float64(pollDelay) * pollDelayMultiplier)
|
||||
if pollDelay > maxPollDelay {
|
||||
pollDelay = maxPollDelay
|
||||
}
|
||||
} else {
|
||||
if !pollTimer.Stop() {
|
||||
<-pollTimer.C
|
||||
}
|
||||
pollDelay = initPollDelay
|
||||
}
|
||||
pollTimer.Reset(pollDelay)
|
||||
send := func(p []byte) {
|
||||
domain := c.domains[mrand.Intn(len(c.domains))]
|
||||
qtype := domain.types[mrand.Intn(len(domain.types))]
|
||||
|
||||
if c.closed {
|
||||
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
|
||||
}
|
||||
|
||||
cur := c.resolverIdx
|
||||
curSend := c.resolverSend[c.resolverAddrs[cur].String()].Add(1)
|
||||
_, _ = c.PacketConn.WriteTo(p.p, c.resolverAddrs[cur])
|
||||
for {
|
||||
c.resolverIdx += 1
|
||||
c.resolverIdx %= uint32(len(c.resolverAddrs))
|
||||
if c.resolverIdx == cur {
|
||||
break
|
||||
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++
|
||||
}
|
||||
if c.resolverSend[c.resolverAddrs[c.resolverIdx].String()].Load() < curSend {
|
||||
break
|
||||
|
||||
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
|
||||
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 {
|
||||
delay *= pollDelayMultiplier
|
||||
if delay > maxPollDelay {
|
||||
delay = maxPollDelay
|
||||
}
|
||||
timeout = false
|
||||
} else {
|
||||
delay = initPollDelay
|
||||
}
|
||||
ticker.Reset(delay)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsConnClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
||||
packet, ok := <-c.readQueue
|
||||
if !ok {
|
||||
return 0, nil, net.ErrClosed
|
||||
func (c *xdnsClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
||||
packet, ok := <-c.readCh
|
||||
if ok {
|
||||
return copy(p, packet.p), packet.addr, nil
|
||||
}
|
||||
if len(p) < len(packet.p) {
|
||||
errors.LogDebug(context.Background(), packet.addr, " mask read err short buffer ", len(p), " ", len(packet.p))
|
||||
return 0, packet.addr, nil
|
||||
}
|
||||
copy(p, packet.p)
|
||||
return len(packet.p), packet.addr, nil
|
||||
return 0, nil, io.ErrClosedPipe
|
||||
}
|
||||
|
||||
func (c *xdnsConnClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
if c.closed {
|
||||
func (c *xdnsClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed() {
|
||||
return 0, io.ErrClosedPipe
|
||||
}
|
||||
|
||||
idx := c.resolverIdx % uint32(len(c.resolverAddrs))
|
||||
encoded, err := encode(p, c.clientID, c.domains[idx], c.resolverTypes[idx])
|
||||
if err != nil {
|
||||
errors.LogDebug(context.Background(), addr, " xdns wireformat err ", err, " ", len(p))
|
||||
return 0, nil
|
||||
if len(p) == 0 || len(p) > 4096 {
|
||||
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.writeQueue <- &packet{
|
||||
p: encoded,
|
||||
addr: addr,
|
||||
}:
|
||||
return len(p), nil
|
||||
case c.sendCh <- b:
|
||||
default:
|
||||
errors.LogDebug(context.Background(), addr, " mask write err queue full")
|
||||
return 0, nil
|
||||
}
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func (c *xdnsConnClient) Close() error {
|
||||
c.closed = true
|
||||
return c.PacketConn.Close()
|
||||
}
|
||||
|
||||
func encode(p []byte, clientID []byte, domain Name, qtype uint16) ([]byte, error) {
|
||||
var decoded []byte
|
||||
{
|
||||
if len(p) >= 224 {
|
||||
return nil, errors.New("too long")
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
buf.Write(clientID[:])
|
||||
n := numPadding
|
||||
if len(p) == 0 {
|
||||
n = numPaddingForPoll
|
||||
}
|
||||
buf.WriteByte(byte(224 + n))
|
||||
_, _ = io.CopyN(&buf, rand.Reader, int64(n))
|
||||
if len(p) > 0 {
|
||||
buf.WriteByte(byte(len(p)))
|
||||
buf.Write(p)
|
||||
}
|
||||
decoded = buf.Bytes()
|
||||
}
|
||||
|
||||
encoded := make([]byte, base32Encoding.EncodedLen(len(decoded)))
|
||||
base32Encoding.Encode(encoded, decoded)
|
||||
encoded = bytes.ToLower(encoded)
|
||||
labels := chunks(encoded, 63)
|
||||
labels = append(labels, domain...)
|
||||
name, err := NewName(labels)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var id uint16
|
||||
_ = binary.Read(rand.Reader, binary.BigEndian, &id)
|
||||
query := &Message{
|
||||
ID: id,
|
||||
Flags: 0x0100,
|
||||
Question: []Question{
|
||||
{
|
||||
Name: name,
|
||||
Type: qtype,
|
||||
Class: ClassIN,
|
||||
},
|
||||
},
|
||||
Additional: []RR{
|
||||
{
|
||||
Name: Name{},
|
||||
Type: RRTypeOPT,
|
||||
Class: 4096,
|
||||
TTL: 0,
|
||||
Data: []byte{},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
buf, err := query.WireFormat()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return buf, nil
|
||||
}
|
||||
|
||||
func chunks(p []byte, n int) [][]byte {
|
||||
var result [][]byte
|
||||
for len(p) > 0 {
|
||||
sz := len(p)
|
||||
if sz > n {
|
||||
sz = n
|
||||
}
|
||||
result = append(result, p[:sz])
|
||||
p = p[sz:]
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func nextPacket(r *bytes.Reader) ([]byte, error) {
|
||||
var n uint16
|
||||
err := binary.Read(r, binary.BigEndian, &n)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
p := make([]byte, n)
|
||||
_, err = io.ReadFull(r, p)
|
||||
if err == io.EOF {
|
||||
err = io.ErrUnexpectedEOF
|
||||
}
|
||||
return p, err
|
||||
}
|
||||
|
||||
func dnsResponsePayload(resp *Message, domains []Name) []byte {
|
||||
if resp.Flags&0x8000 != 0x8000 {
|
||||
func (c *xdnsClient) Close() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed() {
|
||||
return nil
|
||||
}
|
||||
if resp.Flags&0x000f != RcodeNoError {
|
||||
return nil
|
||||
close(c.closeCh)
|
||||
for i := range c.resolvers {
|
||||
c.resolvers[i].Close()
|
||||
}
|
||||
|
||||
if len(resp.Answer) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, answer := range resp.Answer {
|
||||
var ok bool
|
||||
for _, domain := range domains {
|
||||
_, ok = answer.Name.TrimSuffix(domain)
|
||||
if ok {
|
||||
break
|
||||
}
|
||||
}
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
return decodeResponsePayload(resp.Answer)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *xdnsClient) LocalAddr() net.Addr { return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} }
|
||||
|
||||
func (c *xdnsClient) SetDeadline(t time.Time) error { return errors.New("not support") }
|
||||
|
||||
func (c *xdnsClient) SetReadDeadline(t time.Time) error { return errors.New("not support") }
|
||||
|
||||
func (c *xdnsClient) SetWriteDeadline(t time.Time) error { return errors.New("not support") }
|
||||
|
||||
type ClientID [8]byte
|
||||
|
||||
func NewClientID() ClientID {
|
||||
var id ClientID
|
||||
common.Must2(rand.Read(id[:]))
|
||||
id[0] &= 0xFC
|
||||
return id
|
||||
}
|
||||
|
||||
func ClientIDFromRaw(id [8]byte) ClientID {
|
||||
id[0] &= 0xFC
|
||||
return id
|
||||
}
|
||||
|
||||
func ClientIDFromAddr(addr *net.UDPAddr) ClientID {
|
||||
return ClientID(addr.IP[8:])
|
||||
}
|
||||
|
||||
func (id ClientID) Addr() *net.UDPAddr {
|
||||
var ip [16]byte
|
||||
ip[0] = 0xFD
|
||||
copy(ip[8:], id[:])
|
||||
return &net.UDPAddr{IP: ip[:]}
|
||||
}
|
||||
|
||||
@@ -6,9 +6,9 @@ import (
|
||||
)
|
||||
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewConnClient(c, conn)
|
||||
return NewClient(c, dialer)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewConnServer(c, conn)
|
||||
return NewServer(c, conn)
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
serial "github.com/xtls/xray-core/common/serial"
|
||||
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
|
||||
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
|
||||
reflect "reflect"
|
||||
@@ -21,17 +22,94 @@ const (
|
||||
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
|
||||
)
|
||||
|
||||
type DomainProto struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Name string `protobuf:"bytes,1,opt,name=name,proto3" json:"name,omitempty"`
|
||||
LenLimit int32 `protobuf:"varint,2,opt,name=len_limit,json=lenLimit,proto3" json:"len_limit,omitempty"`
|
||||
LabelLimit int32 `protobuf:"varint,3,opt,name=label_limit,json=labelLimit,proto3" json:"label_limit,omitempty"`
|
||||
Types []int32 `protobuf:"varint,4,rep,packed,name=types,proto3" json:"types,omitempty"`
|
||||
Edns0 int32 `protobuf:"varint,5,opt,name=edns0,proto3" json:"edns0,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *DomainProto) Reset() {
|
||||
*x = DomainProto{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *DomainProto) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*DomainProto) ProtoMessage() {}
|
||||
|
||||
func (x *DomainProto) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_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 DomainProto.ProtoReflect.Descriptor instead.
|
||||
func (*DomainProto) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{0}
|
||||
}
|
||||
|
||||
func (x *DomainProto) GetName() string {
|
||||
if x != nil {
|
||||
return x.Name
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *DomainProto) GetLenLimit() int32 {
|
||||
if x != nil {
|
||||
return x.LenLimit
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *DomainProto) GetLabelLimit() int32 {
|
||||
if x != nil {
|
||||
return x.LabelLimit
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *DomainProto) GetTypes() []int32 {
|
||||
if x != nil {
|
||||
return x.Types
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *DomainProto) GetEdns0() int32 {
|
||||
if x != nil {
|
||||
return x.Edns0
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Domains []string `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"`
|
||||
Resolvers []string `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"`
|
||||
Domains []*DomainProto `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"`
|
||||
Resolvers []*serial.TypedMessage `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"`
|
||||
ExtraPoll int32 `protobuf:"varint,3,opt,name=extra_poll,json=extraPoll,proto3" json:"extra_poll,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *Config) Reset() {
|
||||
*x = Config{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0]
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
@@ -43,7 +121,7 @@ func (x *Config) String() string {
|
||||
func (*Config) ProtoMessage() {}
|
||||
|
||||
func (x *Config) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0]
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||
if x != nil {
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
if ms.LoadMessageInfo() == nil {
|
||||
@@ -56,31 +134,139 @@ func (x *Config) ProtoReflect() protoreflect.Message {
|
||||
|
||||
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
|
||||
func (*Config) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{0}
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{1}
|
||||
}
|
||||
|
||||
func (x *Config) GetDomains() []string {
|
||||
func (x *Config) GetDomains() []*DomainProto {
|
||||
if x != nil {
|
||||
return x.Domains
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *Config) GetResolvers() []string {
|
||||
func (x *Config) GetResolvers() []*serial.TypedMessage {
|
||||
if x != nil {
|
||||
return x.Resolvers
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *Config) GetExtraPoll() int32 {
|
||||
if x != nil {
|
||||
return x.ExtraPoll
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
type TCPResolverProto struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *TCPResolverProto) Reset() {
|
||||
*x = TCPResolverProto{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *TCPResolverProto) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*TCPResolverProto) ProtoMessage() {}
|
||||
|
||||
func (x *TCPResolverProto) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
|
||||
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 TCPResolverProto.ProtoReflect.Descriptor instead.
|
||||
func (*TCPResolverProto) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{2}
|
||||
}
|
||||
|
||||
func (x *TCPResolverProto) GetAddr() string {
|
||||
if x != nil {
|
||||
return x.Addr
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type UDPResolverProto struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *UDPResolverProto) Reset() {
|
||||
*x = UDPResolverProto{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *UDPResolverProto) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*UDPResolverProto) ProtoMessage() {}
|
||||
|
||||
func (x *UDPResolverProto) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3]
|
||||
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 UDPResolverProto.ProtoReflect.Descriptor instead.
|
||||
func (*UDPResolverProto) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{3}
|
||||
}
|
||||
|
||||
func (x *UDPResolverProto) GetAddr() string {
|
||||
if x != nil {
|
||||
return x.Addr
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
var File_transport_internet_finalmask_xdns_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\"@\n" +
|
||||
"\x06Config\x12\x18\n" +
|
||||
"\adomains\x18\x01 \x03(\tR\adomains\x12\x1c\n" +
|
||||
"\tresolvers\x18\x02 \x03(\tR\tresolversB\x94\x01\n" +
|
||||
".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\x1a!common/serial/typed_message.proto\"\x8b\x01\n" +
|
||||
"\vDomainProto\x12\x12\n" +
|
||||
"\x04name\x18\x01 \x01(\tR\x04name\x12\x1b\n" +
|
||||
"\tlen_limit\x18\x02 \x01(\x05R\blenLimit\x12\x1f\n" +
|
||||
"\vlabel_limit\x18\x03 \x01(\x05R\n" +
|
||||
"labelLimit\x12\x14\n" +
|
||||
"\x05types\x18\x04 \x03(\x05R\x05types\x12\x14\n" +
|
||||
"\x05edns0\x18\x05 \x01(\x05R\x05edns0\"\xb6\x01\n" +
|
||||
"\x06Config\x12M\n" +
|
||||
"\adomains\x18\x01 \x03(\v23.xray.transport.internet.finalmask.xdns.DomainProtoR\adomains\x12>\n" +
|
||||
"\tresolvers\x18\x02 \x03(\v2 .xray.common.serial.TypedMessageR\tresolvers\x12\x1d\n" +
|
||||
"\n" +
|
||||
"extra_poll\x18\x03 \x01(\x05R\textraPoll\"&\n" +
|
||||
"\x10TCPResolverProto\x12\x12\n" +
|
||||
"\x04addr\x18\x01 \x01(\tR\x04addr\"&\n" +
|
||||
"\x10UDPResolverProto\x12\x12\n" +
|
||||
"\x04addr\x18\x01 \x01(\tR\x04addrB\x94\x01\n" +
|
||||
"*com.xray.transport.internet.finalmask.xdnsP\x01Z;github.com/xtls/xray-core/transport/internet/finalmask/xdns\xaa\x02&Xray.Transport.Internet.Finalmask.Xdnsb\x06proto3"
|
||||
|
||||
var (
|
||||
@@ -95,16 +281,22 @@ func file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP() []byte {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescData
|
||||
}
|
||||
|
||||
var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
|
||||
var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 4)
|
||||
var file_transport_internet_finalmask_xdns_config_proto_goTypes = []any{
|
||||
(*Config)(nil), // 0: xray.transport.internet.finalmask.xdns.Config
|
||||
(*DomainProto)(nil), // 0: xray.transport.internet.finalmask.xdns.DomainProto
|
||||
(*Config)(nil), // 1: xray.transport.internet.finalmask.xdns.Config
|
||||
(*TCPResolverProto)(nil), // 2: xray.transport.internet.finalmask.xdns.TCPResolverProto
|
||||
(*UDPResolverProto)(nil), // 3: xray.transport.internet.finalmask.xdns.UDPResolverProto
|
||||
(*serial.TypedMessage)(nil), // 4: xray.common.serial.TypedMessage
|
||||
}
|
||||
var file_transport_internet_finalmask_xdns_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
|
||||
0, // 0: xray.transport.internet.finalmask.xdns.Config.domains:type_name -> xray.transport.internet.finalmask.xdns.DomainProto
|
||||
4, // 1: xray.transport.internet.finalmask.xdns.Config.resolvers:type_name -> xray.common.serial.TypedMessage
|
||||
2, // [2:2] is the sub-list for method output_type
|
||||
2, // [2:2] is the sub-list for method input_type
|
||||
2, // [2:2] is the sub-list for extension type_name
|
||||
2, // [2:2] is the sub-list for extension extendee
|
||||
0, // [0:2] is the sub-list for field type_name
|
||||
}
|
||||
|
||||
func init() { file_transport_internet_finalmask_xdns_config_proto_init() }
|
||||
@@ -118,7 +310,7 @@ func file_transport_internet_finalmask_xdns_config_proto_init() {
|
||||
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
||||
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_xdns_config_proto_rawDesc), len(file_transport_internet_finalmask_xdns_config_proto_rawDesc)),
|
||||
NumEnums: 0,
|
||||
NumMessages: 1,
|
||||
NumMessages: 4,
|
||||
NumExtensions: 0,
|
||||
NumServices: 0,
|
||||
},
|
||||
|
||||
@@ -6,7 +6,26 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/xdns
|
||||
option java_package = "com.xray.transport.internet.finalmask.xdns";
|
||||
option java_multiple_files = true;
|
||||
|
||||
import "common/serial/typed_message.proto";
|
||||
|
||||
message DomainProto {
|
||||
string name = 1;
|
||||
int32 len_limit = 2;
|
||||
int32 label_limit = 3;
|
||||
repeated int32 types = 4;
|
||||
int32 edns0 = 5;
|
||||
}
|
||||
|
||||
message Config {
|
||||
repeated string domains = 1;
|
||||
repeated string resolvers = 2;
|
||||
repeated DomainProto domains = 1;
|
||||
repeated xray.common.serial.TypedMessage resolvers = 2;
|
||||
int32 extra_poll = 3;
|
||||
}
|
||||
|
||||
message TCPResolverProto {
|
||||
string addr = 1;
|
||||
}
|
||||
|
||||
message UDPResolverProto {
|
||||
string addr = 1;
|
||||
}
|
||||
@@ -1,581 +0,0 @@
|
||||
// Package dns deals with encoding and decoding DNS wire format.
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// The maximum number of DNS name compression pointers we are willing to follow.
|
||||
// Without something like this, infinite loops are possible.
|
||||
const compressionPointerLimit = 10
|
||||
|
||||
var (
|
||||
// ErrZeroLengthLabel is the error returned for names that contain a
|
||||
// zero-length label, like "example..com".
|
||||
ErrZeroLengthLabel = errors.New("name contains a zero-length label")
|
||||
|
||||
// ErrLabelTooLong is the error returned for labels that are longer than
|
||||
// 63 octets.
|
||||
ErrLabelTooLong = errors.New("name contains a label longer than 63 octets")
|
||||
|
||||
// ErrNameTooLong is the error returned for names whose encoded
|
||||
// representation is longer than 255 octets.
|
||||
ErrNameTooLong = errors.New("name is longer than 255 octets")
|
||||
|
||||
// ErrReservedLabelType is the error returned when reading a label type
|
||||
// prefix whose two most significant bits are not 00 or 11.
|
||||
ErrReservedLabelType = errors.New("reserved label type")
|
||||
|
||||
// ErrTooManyPointers is the error returned when reading a compressed
|
||||
// name that has too many compression pointers.
|
||||
ErrTooManyPointers = errors.New("too many compression pointers")
|
||||
|
||||
// ErrTrailingBytes is the error returned when bytes remain in the parse
|
||||
// buffer after parsing a message.
|
||||
ErrTrailingBytes = errors.New("trailing bytes after message")
|
||||
|
||||
// ErrIntegerOverflow is the error returned when trying to encode an
|
||||
// integer greater than 65535 into a 16-bit field.
|
||||
ErrIntegerOverflow = errors.New("integer overflow")
|
||||
)
|
||||
|
||||
const (
|
||||
// https://tools.ietf.org/html/rfc1035#section-3.2.2
|
||||
RRTypeA = 1
|
||||
// https://tools.ietf.org/html/rfc1035#section-3.2.2
|
||||
RRTypeCNAME = 5
|
||||
// https://tools.ietf.org/html/rfc1035#section-3.2.2
|
||||
RRTypeTXT = 16
|
||||
// https://tools.ietf.org/html/rfc3596#section-2.1
|
||||
RRTypeAAAA = 28
|
||||
// https://tools.ietf.org/html/rfc6891#section-6.1.1
|
||||
RRTypeOPT = 41
|
||||
|
||||
// https://tools.ietf.org/html/rfc1035#section-3.2.4
|
||||
ClassIN = 1
|
||||
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.1
|
||||
RcodeNoError = 0 // a.k.a. NOERROR
|
||||
RcodeFormatError = 1 // a.k.a. FORMERR
|
||||
RcodeNameError = 3 // a.k.a. NXDOMAIN
|
||||
RcodeNotImplemented = 4 // a.k.a. NOTIMPL
|
||||
// https://tools.ietf.org/html/rfc6891#section-9
|
||||
ExtendedRcodeBadVers = 16 // a.k.a. BADVERS
|
||||
)
|
||||
|
||||
// Name represents a domain name, a sequence of labels each of which is 63
|
||||
// octets or less in length.
|
||||
//
|
||||
// https://tools.ietf.org/html/rfc1035#section-3.1
|
||||
type Name [][]byte
|
||||
|
||||
// NewName returns a Name from a slice of labels, after checking the labels for
|
||||
// validity. Does not include a zero-length label at the end of the slice.
|
||||
func NewName(labels [][]byte) (Name, error) {
|
||||
name := Name(labels)
|
||||
// https://tools.ietf.org/html/rfc1035#section-2.3.4
|
||||
// Various objects and parameters in the DNS have size limits.
|
||||
// labels 63 octets or less
|
||||
// names 255 octets or less
|
||||
for _, label := range labels {
|
||||
if len(label) == 0 {
|
||||
return nil, ErrZeroLengthLabel
|
||||
}
|
||||
if len(label) > 63 {
|
||||
return nil, ErrLabelTooLong
|
||||
}
|
||||
}
|
||||
// Check the total length.
|
||||
builder := newMessageBuilder()
|
||||
builder.WriteName(name)
|
||||
if len(builder.Bytes()) > 255 {
|
||||
return nil, ErrNameTooLong
|
||||
}
|
||||
return name, nil
|
||||
}
|
||||
|
||||
// ParseName returns a new Name from a string of labels separated by dots, after
|
||||
// checking the name for validity. A single dot at the end of the string is
|
||||
// ignored.
|
||||
func ParseName(s string) (Name, error) {
|
||||
b := bytes.TrimSuffix([]byte(s), []byte("."))
|
||||
if len(b) == 0 {
|
||||
// bytes.Split(b, ".") would return [""] in this case
|
||||
return NewName([][]byte{})
|
||||
} else {
|
||||
return NewName(bytes.Split(b, []byte(".")))
|
||||
}
|
||||
}
|
||||
|
||||
// String returns a reversible string representation of name. Labels are
|
||||
// separated by dots, and any bytes in a label that are outside the set
|
||||
// [0-9A-Za-z-] are replaced with a \xXX hex escape sequence.
|
||||
func (name Name) String() string {
|
||||
if len(name) == 0 {
|
||||
return "."
|
||||
}
|
||||
|
||||
var buf strings.Builder
|
||||
for i, label := range name {
|
||||
if i > 0 {
|
||||
buf.WriteByte('.')
|
||||
}
|
||||
for _, b := range label {
|
||||
if b == '-' ||
|
||||
('0' <= b && b <= '9') ||
|
||||
('A' <= b && b <= 'Z') ||
|
||||
('a' <= b && b <= 'z') {
|
||||
buf.WriteByte(b)
|
||||
} else {
|
||||
fmt.Fprintf(&buf, "\\x%02x", b)
|
||||
}
|
||||
}
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
// TrimSuffix returns a Name with the given suffix removed, if it was present.
|
||||
// The second return value indicates whether the suffix was present. If the
|
||||
// suffix was not present, the first return value is nil.
|
||||
func (name Name) TrimSuffix(suffix Name) (Name, bool) {
|
||||
if len(name) < len(suffix) {
|
||||
return nil, false
|
||||
}
|
||||
split := len(name) - len(suffix)
|
||||
fore, aft := name[:split], name[split:]
|
||||
for i := 0; i < len(aft); i++ {
|
||||
if !bytes.Equal(bytes.ToLower(aft[i]), bytes.ToLower(suffix[i])) {
|
||||
return nil, false
|
||||
}
|
||||
}
|
||||
return fore, true
|
||||
}
|
||||
|
||||
// Message represents a DNS message.
|
||||
//
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1
|
||||
type Message struct {
|
||||
ID uint16
|
||||
Flags uint16
|
||||
|
||||
Question []Question
|
||||
Answer []RR
|
||||
Authority []RR
|
||||
Additional []RR
|
||||
}
|
||||
|
||||
// Opcode extracts the OPCODE part of the Flags field.
|
||||
//
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.1
|
||||
func (message *Message) Opcode() uint16 {
|
||||
return (message.Flags >> 11) & 0xf
|
||||
}
|
||||
|
||||
// Rcode extracts the RCODE part of the Flags field.
|
||||
//
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.1
|
||||
func (message *Message) Rcode() uint16 {
|
||||
return message.Flags & 0x000f
|
||||
}
|
||||
|
||||
// Question represents an entry in the question section of a message.
|
||||
//
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.2
|
||||
type Question struct {
|
||||
Name Name
|
||||
Type uint16
|
||||
Class uint16
|
||||
}
|
||||
|
||||
// RR represents a resource record.
|
||||
//
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.3
|
||||
type RR struct {
|
||||
Name Name
|
||||
Type uint16
|
||||
Class uint16
|
||||
TTL uint32
|
||||
Data []byte
|
||||
}
|
||||
|
||||
// readName parses a DNS name from r. It leaves r positioned just after the
|
||||
// parsed name.
|
||||
func readName(r io.ReadSeeker) (Name, error) {
|
||||
var labels [][]byte
|
||||
// We limit the number of compression pointers we are willing to follow.
|
||||
numPointers := 0
|
||||
// If we followed any compression pointers, we must finally seek to just
|
||||
// past the first pointer.
|
||||
var seekTo int64
|
||||
loop:
|
||||
for {
|
||||
var labelType byte
|
||||
err := binary.Read(r, binary.BigEndian, &labelType)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
switch labelType & 0xc0 {
|
||||
case 0x00:
|
||||
// This is an ordinary label.
|
||||
// https://tools.ietf.org/html/rfc1035#section-3.1
|
||||
length := int(labelType & 0x3f)
|
||||
if length == 0 {
|
||||
break loop
|
||||
}
|
||||
label := make([]byte, length)
|
||||
_, err := io.ReadFull(r, label)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
labels = append(labels, label)
|
||||
case 0xc0:
|
||||
// This is a compression pointer.
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.4
|
||||
upper := labelType & 0x3f
|
||||
var lower byte
|
||||
err := binary.Read(r, binary.BigEndian, &lower)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
offset := (uint16(upper) << 8) | uint16(lower)
|
||||
|
||||
if numPointers == 0 {
|
||||
// The first time we encounter a pointer,
|
||||
// remember our position so we can seek back to
|
||||
// it when done.
|
||||
seekTo, err = r.Seek(0, io.SeekCurrent)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
numPointers++
|
||||
if numPointers > compressionPointerLimit {
|
||||
return nil, ErrTooManyPointers
|
||||
}
|
||||
|
||||
// Follow the pointer and continue.
|
||||
_, err = r.Seek(int64(offset), io.SeekStart)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
default:
|
||||
// "The 10 and 01 combinations are reserved for future
|
||||
// use."
|
||||
return nil, ErrReservedLabelType
|
||||
}
|
||||
}
|
||||
// If we followed any pointers, then seek back to just after the first
|
||||
// one.
|
||||
if numPointers > 0 {
|
||||
_, err := r.Seek(seekTo, io.SeekStart)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return NewName(labels)
|
||||
}
|
||||
|
||||
// readQuestion parses one entry from the Question section. It leaves r
|
||||
// positioned just after the parsed entry.
|
||||
//
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.2
|
||||
func readQuestion(r io.ReadSeeker) (Question, error) {
|
||||
var question Question
|
||||
var err error
|
||||
question.Name, err = readName(r)
|
||||
if err != nil {
|
||||
return question, err
|
||||
}
|
||||
for _, ptr := range []*uint16{&question.Type, &question.Class} {
|
||||
err := binary.Read(r, binary.BigEndian, ptr)
|
||||
if err != nil {
|
||||
return question, err
|
||||
}
|
||||
}
|
||||
|
||||
return question, nil
|
||||
}
|
||||
|
||||
// readRR parses one resource record. It leaves r positioned just after the
|
||||
// parsed resource record.
|
||||
//
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.3
|
||||
func readRR(r io.ReadSeeker) (RR, error) {
|
||||
var rr RR
|
||||
var err error
|
||||
rr.Name, err = readName(r)
|
||||
if err != nil {
|
||||
return rr, err
|
||||
}
|
||||
for _, ptr := range []*uint16{&rr.Type, &rr.Class} {
|
||||
err := binary.Read(r, binary.BigEndian, ptr)
|
||||
if err != nil {
|
||||
return rr, err
|
||||
}
|
||||
}
|
||||
err = binary.Read(r, binary.BigEndian, &rr.TTL)
|
||||
if err != nil {
|
||||
return rr, err
|
||||
}
|
||||
var rdLength uint16
|
||||
err = binary.Read(r, binary.BigEndian, &rdLength)
|
||||
if err != nil {
|
||||
return rr, err
|
||||
}
|
||||
rr.Data = make([]byte, rdLength)
|
||||
_, err = io.ReadFull(r, rr.Data)
|
||||
if err != nil {
|
||||
return rr, err
|
||||
}
|
||||
|
||||
return rr, nil
|
||||
}
|
||||
|
||||
// readMessage parses a complete DNS message. It leaves r positioned just after
|
||||
// the parsed message.
|
||||
func readMessage(r io.ReadSeeker) (Message, error) {
|
||||
var message Message
|
||||
|
||||
// Header section
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.1
|
||||
var qdCount, anCount, nsCount, arCount uint16
|
||||
for _, ptr := range []*uint16{
|
||||
&message.ID, &message.Flags,
|
||||
&qdCount, &anCount, &nsCount, &arCount,
|
||||
} {
|
||||
err := binary.Read(r, binary.BigEndian, ptr)
|
||||
if err != nil {
|
||||
return message, err
|
||||
}
|
||||
}
|
||||
|
||||
// Question section
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.2
|
||||
for i := 0; i < int(qdCount); i++ {
|
||||
question, err := readQuestion(r)
|
||||
if err != nil {
|
||||
return message, err
|
||||
}
|
||||
message.Question = append(message.Question, question)
|
||||
}
|
||||
|
||||
// Answer, Authority, and Additional sections
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.3
|
||||
for _, rec := range []struct {
|
||||
ptr *[]RR
|
||||
count uint16
|
||||
}{
|
||||
{&message.Answer, anCount},
|
||||
{&message.Authority, nsCount},
|
||||
{&message.Additional, arCount},
|
||||
} {
|
||||
for i := 0; i < int(rec.count); i++ {
|
||||
rr, err := readRR(r)
|
||||
if err != nil {
|
||||
return message, err
|
||||
}
|
||||
*rec.ptr = append(*rec.ptr, rr)
|
||||
}
|
||||
}
|
||||
|
||||
return message, nil
|
||||
}
|
||||
|
||||
// MessageFromWireFormat parses a message from buf and returns a Message object.
|
||||
// It returns ErrTrailingBytes if there are bytes remaining in buf after parsing
|
||||
// is done.
|
||||
func MessageFromWireFormat(buf []byte) (Message, error) {
|
||||
r := bytes.NewReader(buf)
|
||||
message, err := readMessage(r)
|
||||
if err == io.EOF {
|
||||
err = io.ErrUnexpectedEOF
|
||||
} else if err == nil {
|
||||
// Check for trailing bytes.
|
||||
_, err = r.ReadByte()
|
||||
if err == io.EOF {
|
||||
err = nil
|
||||
} else if err == nil {
|
||||
err = ErrTrailingBytes
|
||||
}
|
||||
}
|
||||
return message, err
|
||||
}
|
||||
|
||||
// messageBuilder manages the state of serializing a DNS message. Its main
|
||||
// function is to keep track of names already written for the purpose of name
|
||||
// compression.
|
||||
type messageBuilder struct {
|
||||
w bytes.Buffer
|
||||
nameCache map[string]int
|
||||
}
|
||||
|
||||
// newMessageBuilder creates a new messageBuilder with an empty name cache.
|
||||
func newMessageBuilder() *messageBuilder {
|
||||
return &messageBuilder{
|
||||
nameCache: make(map[string]int),
|
||||
}
|
||||
}
|
||||
|
||||
// Bytes returns the serialized DNS message as a slice of bytes.
|
||||
func (builder *messageBuilder) Bytes() []byte {
|
||||
return builder.w.Bytes()
|
||||
}
|
||||
|
||||
// WriteName appends name to the in-progress messageBuilder, employing
|
||||
// compression pointers to previously written names if possible.
|
||||
func (builder *messageBuilder) WriteName(name Name) {
|
||||
// https://tools.ietf.org/html/rfc1035#section-3.1
|
||||
for i := range name {
|
||||
// Has this suffix already been encoded in the message?
|
||||
if ptr, ok := builder.nameCache[name[i:].String()]; ok && ptr&0x3fff == ptr {
|
||||
// If so, we can write a compression pointer.
|
||||
binary.Write(&builder.w, binary.BigEndian, uint16(0xc000|ptr))
|
||||
return
|
||||
}
|
||||
// Not cached; we must encode this label verbatim. Store a cache
|
||||
// entry pointing to the beginning of it.
|
||||
builder.nameCache[name[i:].String()] = builder.w.Len()
|
||||
length := len(name[i])
|
||||
if length == 0 || length > 63 {
|
||||
panic(length)
|
||||
}
|
||||
builder.w.WriteByte(byte(length))
|
||||
builder.w.Write(name[i])
|
||||
}
|
||||
builder.w.WriteByte(0)
|
||||
}
|
||||
|
||||
// WriteQuestion appends a Question section entry to the in-progress
|
||||
// messageBuilder.
|
||||
func (builder *messageBuilder) WriteQuestion(question *Question) {
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.2
|
||||
builder.WriteName(question.Name)
|
||||
binary.Write(&builder.w, binary.BigEndian, question.Type)
|
||||
binary.Write(&builder.w, binary.BigEndian, question.Class)
|
||||
}
|
||||
|
||||
// WriteRR appends a resource record to the in-progress messageBuilder. It
|
||||
// returns ErrIntegerOverflow if the length of rr.Data does not fit in 16 bits.
|
||||
func (builder *messageBuilder) WriteRR(rr *RR) error {
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.3
|
||||
builder.WriteName(rr.Name)
|
||||
binary.Write(&builder.w, binary.BigEndian, rr.Type)
|
||||
binary.Write(&builder.w, binary.BigEndian, rr.Class)
|
||||
binary.Write(&builder.w, binary.BigEndian, rr.TTL)
|
||||
rdLength := uint16(len(rr.Data))
|
||||
if int(rdLength) != len(rr.Data) {
|
||||
return ErrIntegerOverflow
|
||||
}
|
||||
binary.Write(&builder.w, binary.BigEndian, rdLength)
|
||||
builder.w.Write(rr.Data)
|
||||
return nil
|
||||
}
|
||||
|
||||
// WriteMessage appends a complete DNS message to the in-progress
|
||||
// messageBuilder. It returns ErrIntegerOverflow if the number of entries in any
|
||||
// section, or the length of the data in any resource record, does not fit in 16
|
||||
// bits.
|
||||
func (builder *messageBuilder) WriteMessage(message *Message) error {
|
||||
// Header section
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.1
|
||||
binary.Write(&builder.w, binary.BigEndian, message.ID)
|
||||
binary.Write(&builder.w, binary.BigEndian, message.Flags)
|
||||
for _, count := range []int{
|
||||
len(message.Question),
|
||||
len(message.Answer),
|
||||
len(message.Authority),
|
||||
len(message.Additional),
|
||||
} {
|
||||
count16 := uint16(count)
|
||||
if int(count16) != count {
|
||||
return ErrIntegerOverflow
|
||||
}
|
||||
binary.Write(&builder.w, binary.BigEndian, count16)
|
||||
}
|
||||
|
||||
// Question section
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.2
|
||||
for _, question := range message.Question {
|
||||
builder.WriteQuestion(&question)
|
||||
}
|
||||
|
||||
// Answer, Authority, and Additional sections
|
||||
// https://tools.ietf.org/html/rfc1035#section-4.1.3
|
||||
for _, rrs := range [][]RR{message.Answer, message.Authority, message.Additional} {
|
||||
for _, rr := range rrs {
|
||||
err := builder.WriteRR(&rr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// WireFormat encodes a Message as a slice of bytes in DNS wire format. It
|
||||
// returns ErrIntegerOverflow if the number of entries in any section, or the
|
||||
// length of the data in any resource record, does not fit in 16 bits.
|
||||
func (message *Message) WireFormat() ([]byte, error) {
|
||||
builder := newMessageBuilder()
|
||||
err := builder.WriteMessage(message)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return builder.Bytes(), nil
|
||||
}
|
||||
|
||||
// DecodeRDataTXT decodes TXT-DATA (as found in the RDATA for a resource record
|
||||
// with TYPE=TXT) as a raw byte slice, by concatenating all the
|
||||
// <character-string>s it contains.
|
||||
//
|
||||
// https://tools.ietf.org/html/rfc1035#section-3.3.14
|
||||
func DecodeRDataTXT(p []byte) ([]byte, error) {
|
||||
var buf bytes.Buffer
|
||||
for {
|
||||
if len(p) == 0 {
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
n := int(p[0])
|
||||
p = p[1:]
|
||||
if len(p) < n {
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
buf.Write(p[:n])
|
||||
p = p[n:]
|
||||
if len(p) == 0 {
|
||||
break
|
||||
}
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
// EncodeRDataTXT encodes a slice of bytes as TXT-DATA, as appropriate for the
|
||||
// RDATA of a resource record with TYPE=TXT. No length restriction is enforced
|
||||
// here; that must be checked at a higher level.
|
||||
//
|
||||
// https://tools.ietf.org/html/rfc1035#section-3.3.14
|
||||
func EncodeRDataTXT(p []byte) []byte {
|
||||
// https://tools.ietf.org/html/rfc1035#section-3.3
|
||||
// https://tools.ietf.org/html/rfc1035#section-3.3.14
|
||||
// TXT data is a sequence of one or more <character-string>s, where
|
||||
// <character-string> is a length octet followed by that number of
|
||||
// octets.
|
||||
var buf bytes.Buffer
|
||||
for len(p) > 255 {
|
||||
buf.WriteByte(255)
|
||||
buf.Write(p[:255])
|
||||
p = p[255:]
|
||||
}
|
||||
// Must write here, even if len(p) == 0, because it's "*one or more*
|
||||
// <character-string>s".
|
||||
buf.WriteByte(byte(len(p)))
|
||||
buf.Write(p)
|
||||
return buf.Bytes()
|
||||
}
|
||||
@@ -1,953 +0,0 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func namesEqual(a, b Name) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
for i := 0; i < len(a); i++ {
|
||||
if !bytes.Equal(a[i], b[i]) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func TestName(t *testing.T) {
|
||||
for _, test := range []struct {
|
||||
labels [][]byte
|
||||
err error
|
||||
s string
|
||||
}{
|
||||
{[][]byte{}, nil, "."},
|
||||
{[][]byte{[]byte("test")}, nil, "test"},
|
||||
{[][]byte{[]byte("a"), []byte("b"), []byte("c")}, nil, "a.b.c"},
|
||||
|
||||
{[][]byte{{}}, ErrZeroLengthLabel, ""},
|
||||
{[][]byte{[]byte("a"), {}, []byte("c")}, ErrZeroLengthLabel, ""},
|
||||
|
||||
// 63 octets.
|
||||
{
|
||||
[][]byte{[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE")},
|
||||
nil,
|
||||
"0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE",
|
||||
},
|
||||
// 64 octets.
|
||||
{[][]byte{[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDEF")}, ErrLabelTooLong, ""},
|
||||
|
||||
// 64+64+64+62 octets.
|
||||
{
|
||||
[][]byte{
|
||||
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
|
||||
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
|
||||
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
|
||||
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABC"),
|
||||
},
|
||||
nil,
|
||||
"0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABC",
|
||||
},
|
||||
// 64+64+64+63 octets.
|
||||
{[][]byte{
|
||||
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
|
||||
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
|
||||
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
|
||||
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCD"),
|
||||
}, ErrNameTooLong, ""},
|
||||
// 127 one-octet labels.
|
||||
{
|
||||
[][]byte{
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'a'},
|
||||
{'b'},
|
||||
{'c'},
|
||||
{'d'},
|
||||
{'e'},
|
||||
{'f'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'A'},
|
||||
{'B'},
|
||||
{'C'},
|
||||
{'D'},
|
||||
{'E'},
|
||||
{'F'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'a'},
|
||||
{'b'},
|
||||
{'c'},
|
||||
{'d'},
|
||||
{'e'},
|
||||
{'f'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'A'},
|
||||
{'B'},
|
||||
{'C'},
|
||||
{'D'},
|
||||
{'E'},
|
||||
{'F'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'a'},
|
||||
{'b'},
|
||||
{'c'},
|
||||
{'d'},
|
||||
{'e'},
|
||||
{'f'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'A'},
|
||||
{'B'},
|
||||
{'C'},
|
||||
{'D'},
|
||||
{'E'},
|
||||
{'F'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'a'},
|
||||
{'b'},
|
||||
{'c'},
|
||||
{'d'},
|
||||
{'e'},
|
||||
{'f'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'A'},
|
||||
{'B'},
|
||||
{'C'},
|
||||
{'D'},
|
||||
{'E'},
|
||||
},
|
||||
nil,
|
||||
"0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E",
|
||||
},
|
||||
// 128 one-octet labels.
|
||||
{[][]byte{
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'a'},
|
||||
{'b'},
|
||||
{'c'},
|
||||
{'d'},
|
||||
{'e'},
|
||||
{'f'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'A'},
|
||||
{'B'},
|
||||
{'C'},
|
||||
{'D'},
|
||||
{'E'},
|
||||
{'F'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'a'},
|
||||
{'b'},
|
||||
{'c'},
|
||||
{'d'},
|
||||
{'e'},
|
||||
{'f'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'A'},
|
||||
{'B'},
|
||||
{'C'},
|
||||
{'D'},
|
||||
{'E'},
|
||||
{'F'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'a'},
|
||||
{'b'},
|
||||
{'c'},
|
||||
{'d'},
|
||||
{'e'},
|
||||
{'f'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'A'},
|
||||
{'B'},
|
||||
{'C'},
|
||||
{'D'},
|
||||
{'E'},
|
||||
{'F'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'a'},
|
||||
{'b'},
|
||||
{'c'},
|
||||
{'d'},
|
||||
{'e'},
|
||||
{'f'},
|
||||
{'0'},
|
||||
{'1'},
|
||||
{'2'},
|
||||
{'3'},
|
||||
{'4'},
|
||||
{'5'},
|
||||
{'6'},
|
||||
{'7'},
|
||||
{'8'},
|
||||
{'9'},
|
||||
{'A'},
|
||||
{'B'},
|
||||
{'C'},
|
||||
{'D'},
|
||||
{'E'},
|
||||
{'F'},
|
||||
}, ErrNameTooLong, ""},
|
||||
} {
|
||||
// Test that NewName returns proper error codes, and otherwise
|
||||
// returns an equal slice of labels.
|
||||
name, err := NewName(test.labels)
|
||||
if err != test.err || (err == nil && !namesEqual(name, test.labels)) {
|
||||
t.Errorf("%+q returned (%+q, %v), expected (%+q, %v)",
|
||||
test.labels, name, err, test.labels, test.err)
|
||||
continue
|
||||
}
|
||||
if test.err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
// Test that the string version of the name comes out as
|
||||
// expected.
|
||||
s := name.String()
|
||||
if s != test.s {
|
||||
t.Errorf("%+q became string %+q, expected %+q", test.labels, s, test.s)
|
||||
continue
|
||||
}
|
||||
|
||||
// Test that parsing from a string back to a Name results in the
|
||||
// original slice of labels.
|
||||
name, err = ParseName(s)
|
||||
if err != nil || !namesEqual(name, test.labels) {
|
||||
t.Errorf("%+q parsing %+q returned (%+q, %v), expected (%+q, %v)",
|
||||
test.labels, s, name, err, test.labels, nil)
|
||||
continue
|
||||
}
|
||||
// A trailing dot should be ignored.
|
||||
if !strings.HasSuffix(s, ".") {
|
||||
dotName, dotErr := ParseName(s + ".")
|
||||
if dotErr != err || !namesEqual(dotName, name) {
|
||||
t.Errorf("%+q parsing %+q returned (%+q, %v), expected (%+q, %v)",
|
||||
test.labels, s+".", dotName, dotErr, name, err)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseName(t *testing.T) {
|
||||
for _, test := range []struct {
|
||||
s string
|
||||
name Name
|
||||
err error
|
||||
}{
|
||||
// This case can't be tested by TestName above because String
|
||||
// will never produce "" (it produces "." instead).
|
||||
{"", [][]byte{}, nil},
|
||||
} {
|
||||
name, err := ParseName(test.s)
|
||||
if err != test.err || (err == nil && !namesEqual(name, test.name)) {
|
||||
t.Errorf("%+q returned (%+q, %v), expected (%+q, %v)",
|
||||
test.s, name, err, test.name, test.err)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func unescapeString(s string) ([][]byte, error) {
|
||||
if s == "." {
|
||||
return [][]byte{}, nil
|
||||
}
|
||||
|
||||
var result [][]byte
|
||||
for _, label := range strings.Split(s, ".") {
|
||||
var buf bytes.Buffer
|
||||
i := 0
|
||||
for i < len(label) {
|
||||
switch label[i] {
|
||||
case '\\':
|
||||
if i+3 >= len(label) {
|
||||
return nil, fmt.Errorf("truncated escape sequence at index %v", i)
|
||||
}
|
||||
if label[i+1] != 'x' {
|
||||
return nil, fmt.Errorf("malformed escape sequence at index %v", i)
|
||||
}
|
||||
b, err := strconv.ParseUint(string(label[i+2:i+4]), 16, 8)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("malformed hex sequence at index %v", i+2)
|
||||
}
|
||||
buf.WriteByte(byte(b))
|
||||
i += 4
|
||||
default:
|
||||
buf.WriteByte(label[i])
|
||||
i++
|
||||
}
|
||||
}
|
||||
result = append(result, buf.Bytes())
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func TestNameString(t *testing.T) {
|
||||
for _, test := range []struct {
|
||||
name Name
|
||||
s string
|
||||
}{
|
||||
{[][]byte{}, "."},
|
||||
{[][]byte{[]byte("\x00"), []byte("a.b"), []byte("c\nd\\")}, "\\x00.a\\x2eb.c\\x0ad\\x5c"},
|
||||
{[][]byte{
|
||||
[]byte("\x00\x01\x02\x03\x04\x05\x06\x07\x08\t\n\x0b\x0c\r\x0e\x0f\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f !\"#$%&'()*+,-./0123456789:;<=>"),
|
||||
[]byte("?@ABCDEFGHIJKLMNOPQRSTUVWXYZ[\\]^_`abcdefghijklmnopqrstuvwxyz{|}"),
|
||||
[]byte("~\x7f\x80\x81\x82\x83\x84\x85\x86\x87\x88\x89\x8a\x8b\x8c\x8d\x8e\x8f\x90\x91\x92\x93\x94\x95\x96\x97\x98\x99\x9a\x9b\x9c\x9d\x9e\x9f\xa0\xa1\xa2\xa3\xa4\xa5\xa6\xa7\xa8\xa9\xaa\xab\xac\xad\xae\xaf\xb0\xb1\xb2\xb3\xb4\xb5\xb6\xb7\xb8\xb9\xba\xbb\xbc"),
|
||||
[]byte("\xbd\xbe\xbf\xc0\xc1\xc2\xc3\xc4\xc5\xc6\xc7\xc8\xc9\xca\xcb\xcc\xcd\xce\xcf\xd0\xd1\xd2\xd3\xd4\xd5\xd6\xd7\xd8\xd9\xda\xdb\xdc\xdd\xde\xdf\xe0\xe1\xe2\xe3\xe4\xe5\xe6\xe7\xe8\xe9\xea\xeb\xec\xed\xee\xef\xf0\xf1\xf2\xf3\xf4\xf5\xf6\xf7\xf8\xf9\xfa\xfb"),
|
||||
[]byte("\xfc\xfd\xfe\xff"),
|
||||
}, "\\x00\\x01\\x02\\x03\\x04\\x05\\x06\\x07\\x08\\x09\\x0a\\x0b\\x0c\\x0d\\x0e\\x0f\\x10\\x11\\x12\\x13\\x14\\x15\\x16\\x17\\x18\\x19\\x1a\\x1b\\x1c\\x1d\\x1e\\x1f\\x20\\x21\\x22\\x23\\x24\\x25\\x26\\x27\\x28\\x29\\x2a\\x2b\\x2c-\\x2e\\x2f0123456789\\x3a\\x3b\\x3c\\x3d\\x3e.\\x3f\\x40ABCDEFGHIJKLMNOPQRSTUVWXYZ\\x5b\\x5c\\x5d\\x5e\\x5f\\x60abcdefghijklmnopqrstuvwxyz\\x7b\\x7c\\x7d.\\x7e\\x7f\\x80\\x81\\x82\\x83\\x84\\x85\\x86\\x87\\x88\\x89\\x8a\\x8b\\x8c\\x8d\\x8e\\x8f\\x90\\x91\\x92\\x93\\x94\\x95\\x96\\x97\\x98\\x99\\x9a\\x9b\\x9c\\x9d\\x9e\\x9f\\xa0\\xa1\\xa2\\xa3\\xa4\\xa5\\xa6\\xa7\\xa8\\xa9\\xaa\\xab\\xac\\xad\\xae\\xaf\\xb0\\xb1\\xb2\\xb3\\xb4\\xb5\\xb6\\xb7\\xb8\\xb9\\xba\\xbb\\xbc.\\xbd\\xbe\\xbf\\xc0\\xc1\\xc2\\xc3\\xc4\\xc5\\xc6\\xc7\\xc8\\xc9\\xca\\xcb\\xcc\\xcd\\xce\\xcf\\xd0\\xd1\\xd2\\xd3\\xd4\\xd5\\xd6\\xd7\\xd8\\xd9\\xda\\xdb\\xdc\\xdd\\xde\\xdf\\xe0\\xe1\\xe2\\xe3\\xe4\\xe5\\xe6\\xe7\\xe8\\xe9\\xea\\xeb\\xec\\xed\\xee\\xef\\xf0\\xf1\\xf2\\xf3\\xf4\\xf5\\xf6\\xf7\\xf8\\xf9\\xfa\\xfb.\\xfc\\xfd\\xfe\\xff"},
|
||||
} {
|
||||
s := test.name.String()
|
||||
if s != test.s {
|
||||
t.Errorf("%+q escaped to %+q, expected %+q", test.name, s, test.s)
|
||||
continue
|
||||
}
|
||||
unescaped, err := unescapeString(s)
|
||||
if err != nil {
|
||||
t.Errorf("%+q unescaping %+q resulted in error %v", test.name, s, err)
|
||||
continue
|
||||
}
|
||||
if !namesEqual(Name(unescaped), test.name) {
|
||||
t.Errorf("%+q roundtripped through %+q to %+q", test.name, s, unescaped)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNameTrimSuffix(t *testing.T) {
|
||||
for _, test := range []struct {
|
||||
name, suffix string
|
||||
trimmed string
|
||||
ok bool
|
||||
}{
|
||||
{"", "", ".", true},
|
||||
{".", ".", ".", true},
|
||||
{"abc", "", "abc", true},
|
||||
{"abc", ".", "abc", true},
|
||||
{"", "abc", ".", false},
|
||||
{".", "abc", ".", false},
|
||||
{"example.com", "com", "example", true},
|
||||
{"example.com", "net", ".", false},
|
||||
{"example.com", "example.com", ".", true},
|
||||
{"example.com", "test.com", ".", false},
|
||||
{"example.com", "xample.com", ".", false},
|
||||
{"example.com", "example", ".", false},
|
||||
{"example.com", "COM", "example", true},
|
||||
{"EXAMPLE.COM", "com", "EXAMPLE", true},
|
||||
} {
|
||||
tmp, ok := mustParseName(test.name).TrimSuffix(mustParseName(test.suffix))
|
||||
trimmed := tmp.String()
|
||||
if ok != test.ok || trimmed != test.trimmed {
|
||||
t.Errorf("TrimSuffix %+q %+q returned (%+q, %v), expected (%+q, %v)",
|
||||
test.name, test.suffix, trimmed, ok, test.trimmed, test.ok)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadName(t *testing.T) {
|
||||
// Good tests.
|
||||
for _, test := range []struct {
|
||||
start int64
|
||||
end int64
|
||||
input string
|
||||
s string
|
||||
}{
|
||||
// Empty name.
|
||||
{0, 1, "\x00abcd", "."},
|
||||
// No pointers.
|
||||
{12, 25, "AAAABBBBCCCC\x07example\x03com\x00", "example.com"},
|
||||
// Backward pointer.
|
||||
{25, 31, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x0c", "sub.example.com"},
|
||||
// Forward pointer.
|
||||
{0, 4, "\x01a\xc0\x04\x03bcd\x00", "a.bcd"},
|
||||
// Two backwards pointers.
|
||||
{31, 38, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x0c\x04sub2\xc0\x19", "sub2.sub.example.com"},
|
||||
// Forward then backward pointer.
|
||||
{25, 31, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x1f\x04sub2\xc0\x0c", "sub.sub2.example.com"},
|
||||
// Overlapping codons.
|
||||
{0, 4, "\x01a\xc0\x03bcd\x00", "a.bcd"},
|
||||
// Pointer to empty label.
|
||||
{0, 10, "\x07example\xc0\x0a\x00", "example"},
|
||||
{1, 11, "\x00\x07example\xc0\x00", "example"},
|
||||
// Pointer to pointer to empty label.
|
||||
{0, 10, "\x07example\xc0\x0a\xc0\x0c\x00", "example"},
|
||||
{1, 11, "\x00\x07example\xc0\x0c\xc0\x00", "example"},
|
||||
} {
|
||||
r := bytes.NewReader([]byte(test.input))
|
||||
_, err := r.Seek(test.start, io.SeekStart)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
name, err := readName(r)
|
||||
if err != nil {
|
||||
t.Errorf("%+q returned error %s", test.input, err)
|
||||
continue
|
||||
}
|
||||
s := name.String()
|
||||
if s != test.s {
|
||||
t.Errorf("%+q returned %+q, expected %+q", test.input, s, test.s)
|
||||
continue
|
||||
}
|
||||
cur, _ := r.Seek(0, io.SeekCurrent)
|
||||
if cur != test.end {
|
||||
t.Errorf("%+q left offset %d, expected %d", test.input, cur, test.end)
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
// Bad tests.
|
||||
for _, test := range []struct {
|
||||
start int64
|
||||
input string
|
||||
err error
|
||||
}{
|
||||
{0, "", io.ErrUnexpectedEOF},
|
||||
// Reserved label type.
|
||||
{0, "\x80example", ErrReservedLabelType},
|
||||
// Reserved label type.
|
||||
{0, "\x40example", ErrReservedLabelType},
|
||||
// No Terminating empty label.
|
||||
{0, "\x07example\x03com", io.ErrUnexpectedEOF},
|
||||
// Pointer past end of buffer.
|
||||
{0, "\x07example\xc0\xff", io.ErrUnexpectedEOF},
|
||||
// Pointer to self.
|
||||
{0, "\x07example\x03com\xc0\x0c", ErrTooManyPointers},
|
||||
// Pointer to self with intermediate label.
|
||||
{0, "\x07example\x03com\xc0\x08", ErrTooManyPointers},
|
||||
// Two pointers that point to each other.
|
||||
{0, "\xc0\x02\xc0\x00", ErrTooManyPointers},
|
||||
// Two pointers that point to each other, with intermediate labels.
|
||||
{0, "\x01a\xc0\x04\x01b\xc0\x00", ErrTooManyPointers},
|
||||
// EOF while reading label.
|
||||
{0, "\x0aexample", io.ErrUnexpectedEOF},
|
||||
// EOF before second byte of pointer.
|
||||
{0, "\xc0", io.ErrUnexpectedEOF},
|
||||
{0, "\x07example\xc0", io.ErrUnexpectedEOF},
|
||||
} {
|
||||
r := bytes.NewReader([]byte(test.input))
|
||||
_, err := r.Seek(test.start, io.SeekStart)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
name, err := readName(r)
|
||||
if err == io.EOF {
|
||||
err = io.ErrUnexpectedEOF
|
||||
}
|
||||
if err != test.err {
|
||||
t.Errorf("%+q returned (%+q, %v), expected %v", test.input, name, err, test.err)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func mustParseName(s string) Name {
|
||||
name, err := ParseName(s)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
func questionsEqual(a, b *Question) bool {
|
||||
if !namesEqual(a.Name, b.Name) {
|
||||
return false
|
||||
}
|
||||
if a.Type != b.Type || a.Class != b.Class {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func rrsEqual(a, b *RR) bool {
|
||||
if !namesEqual(a.Name, b.Name) {
|
||||
return false
|
||||
}
|
||||
if a.Type != b.Type || a.Class != b.Class || a.TTL != b.TTL {
|
||||
return false
|
||||
}
|
||||
if !bytes.Equal(a.Data, b.Data) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func messagesEqual(a, b *Message) bool {
|
||||
if a.ID != b.ID || a.Flags != b.Flags {
|
||||
return false
|
||||
}
|
||||
if len(a.Question) != len(b.Question) {
|
||||
return false
|
||||
}
|
||||
for i := 0; i < len(a.Question); i++ {
|
||||
if !questionsEqual(&a.Question[i], &b.Question[i]) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
for _, rec := range []struct{ rrA, rrB []RR }{
|
||||
{a.Answer, b.Answer},
|
||||
{a.Authority, b.Authority},
|
||||
{a.Additional, b.Additional},
|
||||
} {
|
||||
if len(rec.rrA) != len(rec.rrB) {
|
||||
return false
|
||||
}
|
||||
for i := 0; i < len(rec.rrA); i++ {
|
||||
if !rrsEqual(&rec.rrA[i], &rec.rrB[i]) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func TestMessageFromWireFormat(t *testing.T) {
|
||||
for _, test := range []struct {
|
||||
buf string
|
||||
expected Message
|
||||
err error
|
||||
}{
|
||||
{
|
||||
"\x12\x34",
|
||||
Message{},
|
||||
io.ErrUnexpectedEOF,
|
||||
},
|
||||
{
|
||||
"\x12\x34\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01",
|
||||
Message{
|
||||
ID: 0x1234,
|
||||
Flags: 0x0100,
|
||||
Question: []Question{
|
||||
{
|
||||
Name: mustParseName("www.example.com"),
|
||||
Type: 1,
|
||||
Class: 1,
|
||||
},
|
||||
},
|
||||
Answer: []RR{},
|
||||
Authority: []RR{},
|
||||
Additional: []RR{},
|
||||
},
|
||||
nil,
|
||||
},
|
||||
{
|
||||
"\x12\x34\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01X",
|
||||
Message{},
|
||||
ErrTrailingBytes,
|
||||
},
|
||||
{
|
||||
"\x12\x34\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01\x03www\x07example\x03com\x00\x00\x01\x00\x01\x00\x00\x00\x80\x00\x04\xc0\x00\x02\x01",
|
||||
Message{
|
||||
ID: 0x1234,
|
||||
Flags: 0x8180,
|
||||
Question: []Question{
|
||||
{
|
||||
Name: mustParseName("www.example.com"),
|
||||
Type: 1,
|
||||
Class: 1,
|
||||
},
|
||||
},
|
||||
Answer: []RR{
|
||||
{
|
||||
Name: mustParseName("www.example.com"),
|
||||
Type: 1,
|
||||
Class: 1,
|
||||
TTL: 128,
|
||||
Data: []byte{192, 0, 2, 1},
|
||||
},
|
||||
},
|
||||
Authority: []RR{},
|
||||
Additional: []RR{},
|
||||
},
|
||||
nil,
|
||||
},
|
||||
} {
|
||||
message, err := MessageFromWireFormat([]byte(test.buf))
|
||||
if err != test.err || (err == nil && !messagesEqual(&message, &test.expected)) {
|
||||
t.Errorf("%+q\nreturned (%+v, %v)\nexpected (%+v, %v)",
|
||||
test.buf, message, err, test.expected, test.err)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessageWireFormatRoundTrip(t *testing.T) {
|
||||
for _, message := range []Message{
|
||||
{
|
||||
ID: 0x1234,
|
||||
Flags: 0x0100,
|
||||
Question: []Question{
|
||||
{
|
||||
Name: mustParseName("www.example.com"),
|
||||
Type: 1,
|
||||
Class: 1,
|
||||
},
|
||||
{
|
||||
Name: mustParseName("www2.example.com"),
|
||||
Type: 2,
|
||||
Class: 2,
|
||||
},
|
||||
},
|
||||
Answer: []RR{
|
||||
{
|
||||
Name: mustParseName("abc"),
|
||||
Type: 2,
|
||||
Class: 3,
|
||||
TTL: 0xffffffff,
|
||||
Data: []byte{1},
|
||||
},
|
||||
{
|
||||
Name: mustParseName("xyz"),
|
||||
Type: 2,
|
||||
Class: 3,
|
||||
TTL: 255,
|
||||
Data: []byte{},
|
||||
},
|
||||
},
|
||||
Authority: []RR{
|
||||
{
|
||||
Name: mustParseName("."),
|
||||
Type: 65535,
|
||||
Class: 65535,
|
||||
TTL: 0,
|
||||
Data: []byte("XXXXXXXXXXXXXXXXXXX"),
|
||||
},
|
||||
},
|
||||
Additional: []RR{},
|
||||
},
|
||||
} {
|
||||
buf, err := message.WireFormat()
|
||||
if err != nil {
|
||||
t.Errorf("%+v cannot make wire format: %v", message, err)
|
||||
continue
|
||||
}
|
||||
message2, err := MessageFromWireFormat(buf)
|
||||
if err != nil {
|
||||
t.Errorf("%+q cannot parse wire format: %v", buf, err)
|
||||
continue
|
||||
}
|
||||
if !messagesEqual(&message, &message2) {
|
||||
t.Errorf("messages unequal\nbefore: %+v\n after: %+v", message, message2)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeRDataTXT(t *testing.T) {
|
||||
for _, test := range []struct {
|
||||
p []byte
|
||||
decoded []byte
|
||||
err error
|
||||
}{
|
||||
{[]byte{}, nil, io.ErrUnexpectedEOF},
|
||||
{[]byte("\x00"), []byte{}, nil},
|
||||
{[]byte("\x01"), nil, io.ErrUnexpectedEOF},
|
||||
} {
|
||||
decoded, err := DecodeRDataTXT(test.p)
|
||||
if err != test.err || (err == nil && !bytes.Equal(decoded, test.decoded)) {
|
||||
t.Errorf("%+q\nreturned (%+q, %v)\nexpected (%+q, %v)",
|
||||
test.p, decoded, err, test.decoded, test.err)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestEncodeRDataTXT(t *testing.T) {
|
||||
// Encoding 0 bytes needs to return at least a single length octet of
|
||||
// zero, not an empty slice.
|
||||
p := make([]byte, 0)
|
||||
encoded := EncodeRDataTXT(p)
|
||||
if len(encoded) < 0 {
|
||||
t.Errorf("EncodeRDataTXT(%v) returned %v", p, encoded)
|
||||
}
|
||||
|
||||
// 255 bytes should be able to be encoded into 256 bytes.
|
||||
p = make([]byte, 255)
|
||||
encoded = EncodeRDataTXT(p)
|
||||
if len(encoded) > 256 {
|
||||
t.Errorf("EncodeRDataTXT(%d bytes) returned %d bytes", len(p), len(encoded))
|
||||
}
|
||||
|
||||
fmt.Println(EncodeRDataTXT(nil))
|
||||
fmt.Println(computeMaxEncodedPayload(maxUDPPayload))
|
||||
}
|
||||
|
||||
func TestRDataTXTRoundTrip(t *testing.T) {
|
||||
for _, p := range [][]byte{
|
||||
{},
|
||||
[]byte("\x00"),
|
||||
{
|
||||
0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f,
|
||||
0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e, 0x1f,
|
||||
0x20, 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f,
|
||||
0x30, 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f,
|
||||
0x40, 0x41, 0x42, 0x43, 0x44, 0x45, 0x46, 0x47, 0x48, 0x49, 0x4a, 0x4b, 0x4c, 0x4d, 0x4e, 0x4f,
|
||||
0x50, 0x51, 0x52, 0x53, 0x54, 0x55, 0x56, 0x57, 0x58, 0x59, 0x5a, 0x5b, 0x5c, 0x5d, 0x5e, 0x5f,
|
||||
0x60, 0x61, 0x62, 0x63, 0x64, 0x65, 0x66, 0x67, 0x68, 0x69, 0x6a, 0x6b, 0x6c, 0x6d, 0x6e, 0x6f,
|
||||
0x70, 0x71, 0x72, 0x73, 0x74, 0x75, 0x76, 0x77, 0x78, 0x79, 0x7a, 0x7b, 0x7c, 0x7d, 0x7e, 0x7f,
|
||||
0x80, 0x81, 0x82, 0x83, 0x84, 0x85, 0x86, 0x87, 0x88, 0x89, 0x8a, 0x8b, 0x8c, 0x8d, 0x8e, 0x8f,
|
||||
0x90, 0x91, 0x92, 0x93, 0x94, 0x95, 0x96, 0x97, 0x98, 0x99, 0x9a, 0x9b, 0x9c, 0x9d, 0x9e, 0x9f,
|
||||
0xa0, 0xa1, 0xa2, 0xa3, 0xa4, 0xa5, 0xa6, 0xa7, 0xa8, 0xa9, 0xaa, 0xab, 0xac, 0xad, 0xae, 0xaf,
|
||||
0xb0, 0xb1, 0xb2, 0xb3, 0xb4, 0xb5, 0xb6, 0xb7, 0xb8, 0xb9, 0xba, 0xbb, 0xbc, 0xbd, 0xbe, 0xbf,
|
||||
0xc0, 0xc1, 0xc2, 0xc3, 0xc4, 0xc5, 0xc6, 0xc7, 0xc8, 0xc9, 0xca, 0xcb, 0xcc, 0xcd, 0xce, 0xcf,
|
||||
0xd0, 0xd1, 0xd2, 0xd3, 0xd4, 0xd5, 0xd6, 0xd7, 0xd8, 0xd9, 0xda, 0xdb, 0xdc, 0xdd, 0xde, 0xdf,
|
||||
0xe0, 0xe1, 0xe2, 0xe3, 0xe4, 0xe5, 0xe6, 0xe7, 0xe8, 0xe9, 0xea, 0xeb, 0xec, 0xed, 0xee, 0xef,
|
||||
0xf0, 0xf1, 0xf2, 0xf3, 0xf4, 0xf5, 0xf6, 0xf7, 0xf8, 0xf9, 0xfa, 0xfb, 0xfc, 0xfd, 0xfe, 0xff,
|
||||
},
|
||||
} {
|
||||
rdata := EncodeRDataTXT(p)
|
||||
decoded, err := DecodeRDataTXT(rdata)
|
||||
if err != nil || !bytes.Equal(decoded, p) {
|
||||
t.Errorf("%+q returned (%+q, %v)", p, decoded, err)
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIPAnswerPayloadRoundTrip(t *testing.T) {
|
||||
for _, rrType := range []uint16{RRTypeA, RRTypeAAAA} {
|
||||
for _, payload := range [][]byte{
|
||||
{},
|
||||
{0x01},
|
||||
[]byte("hello world"),
|
||||
bytes.Repeat([]byte{0xab}, payloadChunkSizeForType(rrType)*3+1),
|
||||
} {
|
||||
question := Question{
|
||||
Name: mustParseName("example.com"),
|
||||
Type: rrType,
|
||||
Class: ClassIN,
|
||||
}
|
||||
answers, err := answersForPayload(question, responseTTL, payload)
|
||||
if err != nil {
|
||||
t.Fatalf("answersForPayload(%d) err = %v", rrType, err)
|
||||
}
|
||||
|
||||
if len(answers) > 1 {
|
||||
answers[0], answers[len(answers)-1] = answers[len(answers)-1], answers[0]
|
||||
}
|
||||
|
||||
decoded := decodeResponsePayload(answers)
|
||||
if !bytes.Equal(decoded, payload) {
|
||||
t.Fatalf("rrType=%d decoded %x want %x", rrType, decoded, payload)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseResolver(t *testing.T) {
|
||||
tests := []struct {
|
||||
resolver string
|
||||
rrType uint16
|
||||
}{
|
||||
{"example.com+udp://1.1.1.1:53", RRTypeTXT},
|
||||
{"example.com:txt+udp://1.1.1.1:53", RRTypeTXT},
|
||||
{"example.com:a+udp://1.1.1.1:53", RRTypeA},
|
||||
{"example.com:aaaa+udp://1.1.1.1:53", RRTypeAAAA},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
domain, server, rrType, err := parseResolver(test.resolver)
|
||||
if err != nil {
|
||||
t.Fatalf("parseResolver(%q) err = %v", test.resolver, err)
|
||||
}
|
||||
if domain.String() != "example.com" || server != "1.1.1.1:53" || rrType != test.rrType {
|
||||
t.Fatalf("parseResolver(%q) = (%q, %q, %d)", test.resolver, domain.String(), server, rrType)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseDomainSpec(t *testing.T) {
|
||||
tests := []struct {
|
||||
spec string
|
||||
def string
|
||||
rrType uint16
|
||||
wantErr bool
|
||||
}{
|
||||
{"example.com", "", 0, false},
|
||||
{"example.com", "txt", RRTypeTXT, false},
|
||||
{"example.com:a", "", RRTypeA, false},
|
||||
{"example.com:aaaa", "", RRTypeAAAA, false},
|
||||
{"example.com:doh", "", 0, true},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
got, err := parseDomainSpec(test.spec, test.def)
|
||||
if test.wantErr {
|
||||
if err == nil {
|
||||
t.Fatalf("parseDomainSpec(%q, %q) err = nil", test.spec, test.def)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("parseDomainSpec(%q, %q) err = %v", test.spec, test.def, err)
|
||||
}
|
||||
if got.name.String() != "example.com" || got.rrType != test.rrType {
|
||||
t.Fatalf("parseDomainSpec(%q, %q) = (%q, %d)", test.spec, test.def, got.name.String(), got.rrType)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestResponseForMethodRestriction(t *testing.T) {
|
||||
query := &Message{
|
||||
ID: 1,
|
||||
Flags: 0x0100,
|
||||
Question: []Question{{
|
||||
Name: mustParseName("abc.example.com"),
|
||||
Type: RRTypeTXT,
|
||||
Class: ClassIN,
|
||||
}},
|
||||
Additional: []RR{{
|
||||
Name: Name{},
|
||||
Type: RRTypeOPT,
|
||||
Class: 4096,
|
||||
}},
|
||||
}
|
||||
|
||||
resp, _ := responseFor(query, []domainSpec{{name: mustParseName("example.com"), rrType: RRTypeA}})
|
||||
if resp == nil || resp.Rcode() != RcodeNameError {
|
||||
t.Fatalf("responseFor method restriction rcode = %v", resp)
|
||||
}
|
||||
|
||||
resp, _ = responseFor(query, []domainSpec{{name: mustParseName("example.com")}})
|
||||
if resp == nil || resp.Rcode() != RcodeNoError {
|
||||
t.Fatalf("responseFor unrestricted rcode = %v", resp)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,215 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"encoding/base32"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/net/dns/dnsmessage"
|
||||
"golang.org/x/net/idna"
|
||||
)
|
||||
|
||||
func Lower(c byte) byte {
|
||||
if c >= 'A' && c <= 'Z' {
|
||||
return c + ('a' - 'A')
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
func ToUpper(b []byte) {
|
||||
for i, c := range b {
|
||||
if c >= 'a' && c <= 'z' {
|
||||
b[i] = c - 'a' + 'A'
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func ToLower(b []byte) {
|
||||
for i, c := range b {
|
||||
if c >= 'A' && c <= 'Z' {
|
||||
b[i] = c - 'A' + 'a'
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func NewTable() ([256]int, [256]int) {
|
||||
var t, t_ [256]int
|
||||
for i := range t {
|
||||
t[i] = base32Encoding.DecodedLen(i)
|
||||
}
|
||||
for i := range t_ {
|
||||
t_[i] = base32Encoding.EncodedLen(i)
|
||||
}
|
||||
return t, t_
|
||||
}
|
||||
|
||||
const (
|
||||
TypeA uint16 = 1
|
||||
TypeCNAME uint16 = 5
|
||||
TypeTXT uint16 = 16
|
||||
TypeAAAA uint16 = 28
|
||||
)
|
||||
|
||||
var (
|
||||
base32Encoding = base32.StdEncoding.WithPadding(base32.NoPadding)
|
||||
table, table_ = NewTable()
|
||||
TypeMap = map[uint16]byte{
|
||||
TypeA: 0,
|
||||
TypeCNAME: 1,
|
||||
TypeTXT: 2,
|
||||
TypeAAAA: 3,
|
||||
}
|
||||
TypeMap_ = map[byte]uint16{
|
||||
0: TypeA,
|
||||
1: TypeCNAME,
|
||||
2: TypeTXT,
|
||||
3: TypeAAAA,
|
||||
}
|
||||
)
|
||||
|
||||
type Domain struct {
|
||||
name dnsmessage.Name
|
||||
lenLimit int
|
||||
labelLimit int
|
||||
types []uint16
|
||||
edns0 uint16
|
||||
|
||||
cap int
|
||||
lenMax int
|
||||
}
|
||||
|
||||
func NewDomain(domain string, lenLimit int, labelLimit int, types []uint16, edns0 uint16) (*Domain, error) {
|
||||
if strings.Contains(domain, "..") {
|
||||
return nil, errors.New("invalid domain")
|
||||
}
|
||||
if lenLimit < 0 || lenLimit > 255 {
|
||||
return nil, errors.New("lenLimit < 0 || lenLimit > 255")
|
||||
}
|
||||
if labelLimit < 0 || labelLimit > 63 {
|
||||
return nil, errors.New("labelLimit < 0 || labelLimit > 63")
|
||||
}
|
||||
if len(types) == 0 {
|
||||
return nil, errors.New("empty types")
|
||||
}
|
||||
for i := range types {
|
||||
switch types[i] {
|
||||
case uint16(dnsmessage.TypeA), uint16(dnsmessage.TypeCNAME), uint16(dnsmessage.TypeTXT), uint16(dnsmessage.TypeAAAA):
|
||||
default:
|
||||
return nil, errors.New("unknown types")
|
||||
}
|
||||
}
|
||||
if edns0 != 0 && (edns0 < 512 || edns0 > 4096) {
|
||||
return nil, errors.New("edns0 != 0 && (edns0 < 512 || edns0 > 4096)")
|
||||
}
|
||||
|
||||
ascii, err := idna.ToASCII(domain)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ascii = strings.Trim(ascii, ".")
|
||||
|
||||
name, err := dnsmessage.NewName(domain + ".")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if lenLimit < int(name.Length)+1 {
|
||||
return nil, errors.New("lenLimit < int(name.Length)+1")
|
||||
}
|
||||
n := (lenLimit - int(name.Length) - 1) / (labelLimit + 1)
|
||||
left := (lenLimit - int(name.Length) - 1) % (labelLimit + 1)
|
||||
total := n * labelLimit
|
||||
if left > 1 {
|
||||
total += left - 1
|
||||
}
|
||||
cap := table[total]
|
||||
if cap < 17 {
|
||||
return nil, errors.New("cap < 17")
|
||||
}
|
||||
total = table_[cap]
|
||||
lenMax := int(name.Length) + 1 + total + total/labelLimit
|
||||
if total%labelLimit > 0 {
|
||||
lenMax += 1
|
||||
}
|
||||
return &Domain{
|
||||
name: name,
|
||||
lenLimit: lenLimit,
|
||||
labelLimit: labelLimit,
|
||||
types: types,
|
||||
edns0: edns0,
|
||||
|
||||
cap: cap,
|
||||
lenMax: lenMax,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (d *Domain) Show() string {
|
||||
return fmt.Sprint(d.name, d.cap)
|
||||
}
|
||||
|
||||
func (d *Domain) IsDomain(name dnsmessage.Name) bool {
|
||||
if d.name.Length >= name.Length {
|
||||
return false
|
||||
}
|
||||
i := d.name.Length
|
||||
j := name.Length
|
||||
for i > 0 {
|
||||
i--
|
||||
j--
|
||||
if Lower(d.name.Data[i]) != Lower(name.Data[j]) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (d *Domain) HasType(qtype uint16) bool {
|
||||
for i := range d.types {
|
||||
if d.types[i] == qtype {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (d *Domain) Encode(data []byte) dnsmessage.Name {
|
||||
var name dnsmessage.Name
|
||||
var encoded [255]byte
|
||||
base32Encoding.Encode(encoded[:], data)
|
||||
ToLower(encoded[:table_[len(data)]])
|
||||
b1 := name.Data[:0]
|
||||
b2 := encoded[:table_[len(data)]]
|
||||
for len(b2) > 0 {
|
||||
size := min(len(b2), d.labelLimit)
|
||||
b1 = append(b1, b2[:size]...)
|
||||
b1 = append(b1, '.')
|
||||
b2 = b2[size:]
|
||||
}
|
||||
b1 = append(b1, d.name.Data[:d.name.Length]...)
|
||||
if len(b1) > 254 {
|
||||
panic("len(b1) > 254")
|
||||
}
|
||||
name.Length = byte(len(b1))
|
||||
return name
|
||||
}
|
||||
|
||||
func (d *Domain) Decode(decoded *[255]byte, name dnsmessage.Name) int {
|
||||
if !d.IsDomain(name) {
|
||||
return 0
|
||||
}
|
||||
var encoded [255]byte
|
||||
b1 := encoded[:0]
|
||||
b2 := name.Data[:name.Length-d.name.Length]
|
||||
for i := range b2 {
|
||||
if b2[i] != '.' {
|
||||
b1 = append(b1, b2[i])
|
||||
}
|
||||
}
|
||||
ToUpper(b1)
|
||||
n, err := base32Encoding.Decode(decoded[:], b1)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return n
|
||||
}
|
||||
@@ -0,0 +1,171 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
fragTTL = 8 * time.Second
|
||||
fragSize = 4096
|
||||
fragClientIDSize = 16384
|
||||
fragCount = 4096
|
||||
)
|
||||
|
||||
type FragKey struct {
|
||||
clientID ClientID
|
||||
fragID byte
|
||||
}
|
||||
|
||||
type FragEntry struct {
|
||||
data [][]byte
|
||||
size int
|
||||
len int
|
||||
total byte
|
||||
deadline time.Time
|
||||
}
|
||||
|
||||
type FragManager struct {
|
||||
m map[FragKey]*FragEntry
|
||||
sizem map[ClientID]int
|
||||
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{}),
|
||||
}
|
||||
go m.gc()
|
||||
return m
|
||||
}
|
||||
|
||||
func (m *FragManager) closed() bool {
|
||||
select {
|
||||
case <-m.ch:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (m *FragManager) removeEntey(k FragKey, e *FragEntry) {
|
||||
m.sizem[k.clientID] -= e.size
|
||||
delete(m.m, k)
|
||||
}
|
||||
|
||||
func (m *FragManager) tryRemove() {
|
||||
if len(m.m) < fragCount {
|
||||
return
|
||||
}
|
||||
var key FragKey
|
||||
var entry *FragEntry
|
||||
first := true
|
||||
for k, e := range m.m {
|
||||
if first || e.deadline.Before(entry.deadline) {
|
||||
key = k
|
||||
entry = e
|
||||
first = false
|
||||
}
|
||||
}
|
||||
m.removeEntey(key, entry)
|
||||
}
|
||||
|
||||
func (m *FragManager) gc() {
|
||||
ticker := time.NewTicker(fragTTL / 2)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-m.ch:
|
||||
return
|
||||
case now := <-ticker.C:
|
||||
m.mu.Lock()
|
||||
for k, e := range m.m {
|
||||
if now.After(e.deadline) {
|
||||
m.removeEntey(k, e)
|
||||
}
|
||||
}
|
||||
m.mu.Unlock()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []byte) int {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.closed() {
|
||||
return 0
|
||||
}
|
||||
|
||||
if fragN < 2 {
|
||||
return 0
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
entry := m.m[key]
|
||||
if entry == nil || now.After(entry.deadline) {
|
||||
if entry == nil {
|
||||
m.tryRemove()
|
||||
} else {
|
||||
m.removeEntey(key, entry)
|
||||
}
|
||||
entry = &FragEntry{
|
||||
data: make([][]byte, fragN),
|
||||
total: fragN,
|
||||
deadline: now.Add(fragTTL),
|
||||
}
|
||||
m.m[key] = entry
|
||||
}
|
||||
|
||||
if fragN != entry.total {
|
||||
return 0
|
||||
}
|
||||
if fragIdx >= entry.total {
|
||||
return 0
|
||||
}
|
||||
if entry.data[fragIdx] != nil {
|
||||
return 0
|
||||
}
|
||||
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)
|
||||
|
||||
entry.data[fragIdx] = cp
|
||||
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
|
||||
}
|
||||
|
||||
out = out[:0]
|
||||
for i := range entry.data {
|
||||
out = append(out, entry.data[i]...)
|
||||
}
|
||||
m.removeEntey(key, entry)
|
||||
return len(out)
|
||||
}
|
||||
|
||||
func (m *FragManager) Close() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.closed() {
|
||||
return
|
||||
}
|
||||
close(m.ch)
|
||||
for k := range m.m {
|
||||
delete(m.m, k)
|
||||
}
|
||||
}
|
||||
@@ -1,226 +0,0 @@
|
||||
package xdns
|
||||
|
||||
import "bytes"
|
||||
|
||||
const ipRecordHeaderSize = 2
|
||||
|
||||
func maxEncodedPayloadForType(rrType uint16) int {
|
||||
switch rrType {
|
||||
case RRTypeA:
|
||||
return maxEncodedPayloadA
|
||||
case RRTypeAAAA:
|
||||
return maxEncodedPayloadAAAA
|
||||
default:
|
||||
return maxEncodedPayloadTXT
|
||||
}
|
||||
}
|
||||
|
||||
func rrDataSizeForType(rrType uint16) int {
|
||||
switch rrType {
|
||||
case RRTypeA:
|
||||
return 4
|
||||
case RRTypeAAAA:
|
||||
return 16
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func payloadChunkSizeForType(rrType uint16) int {
|
||||
size := rrDataSizeForType(rrType)
|
||||
if size <= ipRecordHeaderSize {
|
||||
return 0
|
||||
}
|
||||
return size - ipRecordHeaderSize
|
||||
}
|
||||
|
||||
func answersForPayload(question Question, ttl uint32, payload []byte) ([]RR, error) {
|
||||
switch question.Type {
|
||||
case RRTypeTXT:
|
||||
return []RR{
|
||||
{
|
||||
Name: question.Name,
|
||||
Type: question.Type,
|
||||
Class: question.Class,
|
||||
TTL: ttl,
|
||||
Data: EncodeRDataTXT(payload),
|
||||
},
|
||||
}, nil
|
||||
case RRTypeA, RRTypeAAAA:
|
||||
return ipAnswersForPayload(question, ttl, payload)
|
||||
default:
|
||||
return nil, ErrIntegerOverflow
|
||||
}
|
||||
}
|
||||
|
||||
func ipAnswersForPayload(question Question, ttl uint32, payload []byte) ([]RR, error) {
|
||||
chunkSize := payloadChunkSizeForType(question.Type)
|
||||
rrDataSize := rrDataSizeForType(question.Type)
|
||||
if chunkSize == 0 || rrDataSize == 0 {
|
||||
return nil, ErrIntegerOverflow
|
||||
}
|
||||
|
||||
numRecords := 1
|
||||
if len(payload) > 0 {
|
||||
numRecords = (len(payload) + chunkSize - 1) / chunkSize
|
||||
}
|
||||
if numRecords > 256 {
|
||||
return nil, ErrIntegerOverflow
|
||||
}
|
||||
|
||||
answers := make([]RR, 0, numRecords)
|
||||
for i := 0; i < numRecords; i++ {
|
||||
offset := i * chunkSize
|
||||
n := len(payload) - offset
|
||||
if n < 0 {
|
||||
n = 0
|
||||
}
|
||||
if n > chunkSize {
|
||||
n = chunkSize
|
||||
}
|
||||
|
||||
data := make([]byte, rrDataSize)
|
||||
data[0] = byte(i)
|
||||
data[1] = byte(n)
|
||||
copy(data[ipRecordHeaderSize:], payload[offset:offset+n])
|
||||
|
||||
answers = append(answers, RR{
|
||||
Name: question.Name,
|
||||
Type: question.Type,
|
||||
Class: question.Class,
|
||||
TTL: ttl,
|
||||
Data: data,
|
||||
})
|
||||
}
|
||||
|
||||
return answers, nil
|
||||
}
|
||||
|
||||
func decodeResponsePayload(answers []RR) []byte {
|
||||
if len(answers) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
switch answers[0].Type {
|
||||
case RRTypeTXT:
|
||||
if len(answers) != 1 {
|
||||
return nil
|
||||
}
|
||||
payload, err := DecodeRDataTXT(answers[0].Data)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return payload
|
||||
case RRTypeA, RRTypeAAAA:
|
||||
return decodeIPAnswerPayload(answers, answers[0].Type)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func decodeIPAnswerPayload(answers []RR, rrType uint16) []byte {
|
||||
chunkSize := payloadChunkSizeForType(rrType)
|
||||
rrDataSize := rrDataSizeForType(rrType)
|
||||
if chunkSize == 0 || rrDataSize == 0 || len(answers) > 256 {
|
||||
return nil
|
||||
}
|
||||
|
||||
parts := make([][]byte, len(answers))
|
||||
for _, answer := range answers {
|
||||
if answer.Type != rrType || len(answer.Data) != rrDataSize {
|
||||
return nil
|
||||
}
|
||||
idx := int(answer.Data[0])
|
||||
n := int(answer.Data[1])
|
||||
if idx >= len(answers) || n > chunkSize || parts[idx] != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
part := make([]byte, n)
|
||||
copy(part, answer.Data[ipRecordHeaderSize:ipRecordHeaderSize+n])
|
||||
parts[idx] = part
|
||||
}
|
||||
|
||||
var payload bytes.Buffer
|
||||
for _, part := range parts {
|
||||
if part == nil {
|
||||
return nil
|
||||
}
|
||||
payload.Write(part)
|
||||
}
|
||||
return payload.Bytes()
|
||||
}
|
||||
|
||||
func computeMaxEncodedPayload(limit int) int {
|
||||
return computeMaxEncodedPayloadForType(limit, RRTypeTXT)
|
||||
}
|
||||
|
||||
func computeMaxEncodedPayloadForType(limit int, rrType uint16) int {
|
||||
maxLengthName, err := NewName([][]byte{
|
||||
[]byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"),
|
||||
[]byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"),
|
||||
[]byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"),
|
||||
[]byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"),
|
||||
})
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
{
|
||||
n := 0
|
||||
for _, label := range maxLengthName {
|
||||
n += len(label) + 1
|
||||
}
|
||||
n += 1
|
||||
if n != 255 {
|
||||
panic("computeMaxEncodedPayload n != 255")
|
||||
}
|
||||
}
|
||||
|
||||
queryLimit := uint16(limit)
|
||||
if int(queryLimit) != limit {
|
||||
queryLimit = 0xffff
|
||||
}
|
||||
query := &Message{
|
||||
Question: []Question{
|
||||
{
|
||||
Name: maxLengthName,
|
||||
Type: rrType,
|
||||
Class: ClassIN,
|
||||
},
|
||||
},
|
||||
Additional: []RR{
|
||||
{
|
||||
Name: Name{},
|
||||
Type: RRTypeOPT,
|
||||
Class: queryLimit,
|
||||
TTL: 0,
|
||||
Data: []byte{},
|
||||
},
|
||||
},
|
||||
}
|
||||
resp, _ := responseFor(query, []domainSpec{{name: Name{[]byte{}}}})
|
||||
|
||||
low := 0
|
||||
high := 32768
|
||||
if chunkSize := payloadChunkSizeForType(rrType); chunkSize > 0 {
|
||||
high = 256*chunkSize + 1
|
||||
}
|
||||
for low+1 < high {
|
||||
mid := (low + high) / 2
|
||||
resp.Answer, err = answersForPayload(query.Question[0], responseTTL, make([]byte, mid))
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
buf, err := resp.WireFormat()
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
if len(buf) <= limit {
|
||||
low = mid
|
||||
} else {
|
||||
high = mid
|
||||
}
|
||||
}
|
||||
|
||||
return low
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net"
|
||||
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
type Resolver interface {
|
||||
Addr() *net.UDPAddr
|
||||
Read(p []byte) (int, error)
|
||||
Send(p []byte)
|
||||
Close()
|
||||
}
|
||||
|
||||
func NewResolver(proto *serial.TypedMessage, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
config, err := proto.GetInstance()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch v := config.(type) {
|
||||
case *TCPResolverProto:
|
||||
return NewTCPResolver(v, dialer)
|
||||
case *UDPResolverProto:
|
||||
return NewUDPResolver(v, dialer)
|
||||
default:
|
||||
return nil, errors.New("unknown proto")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"io"
|
||||
"sync"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
type TCPResolver struct {
|
||||
dest net.Destination
|
||||
dialer *finalmask.Dialer
|
||||
|
||||
conn net.Conn
|
||||
tcpAddr *net.TCPAddr
|
||||
udpAddr *net.UDPAddr
|
||||
|
||||
readCh chan []byte
|
||||
closeCh chan struct{}
|
||||
wg sync.WaitGroup
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewTCPResolver(config *TCPResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
dest, err := net.ParseDestination("tcp:" + config.Addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r := &TCPResolver{
|
||||
dest: dest,
|
||||
dialer: dialer,
|
||||
readCh: make(chan []byte),
|
||||
closeCh: make(chan struct{}),
|
||||
}
|
||||
if err := r.dial(); err != nil {
|
||||
r.Close()
|
||||
return nil, err
|
||||
}
|
||||
return r, nil
|
||||
}
|
||||
|
||||
func (r *TCPResolver) closed() bool {
|
||||
select {
|
||||
case <-r.closeCh:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (r *TCPResolver) dial() error {
|
||||
if r.closed() {
|
||||
return errors.New("closed")
|
||||
}
|
||||
if r.conn != nil {
|
||||
return nil
|
||||
}
|
||||
conn, err := r.dialer.DialTCP(r.dest)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.conn = conn
|
||||
r.tcpAddr = conn.RemoteAddr().(*net.TCPAddr)
|
||||
r.udpAddr = &net.UDPAddr{IP: r.tcpAddr.IP, Port: r.tcpAddr.Port}
|
||||
r.wg.Add(1)
|
||||
go r.recv(conn)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *TCPResolver) recv(conn net.Conn) {
|
||||
defer r.wg.Done()
|
||||
|
||||
var buf [4096]byte
|
||||
for {
|
||||
_, err := io.ReadFull(conn, buf[:2])
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
n := binary.BigEndian.Uint16(buf[:2])
|
||||
if n == 0 || n > 4096 {
|
||||
io.CopyN(io.Discard, conn, int64(n))
|
||||
continue
|
||||
}
|
||||
_, err = io.ReadFull(conn, buf[:n])
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
p := pool4K.Get().([]byte)
|
||||
copy(p, buf[:n])
|
||||
select {
|
||||
case <-r.closeCh:
|
||||
pool4K.Put(p[:cap(p)])
|
||||
case r.readCh <- p[:n]:
|
||||
}
|
||||
}
|
||||
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
_ = conn.Close()
|
||||
r.conn = nil
|
||||
}
|
||||
|
||||
func (r *TCPResolver) Addr() *net.UDPAddr {
|
||||
return r.udpAddr
|
||||
}
|
||||
|
||||
func (r *TCPResolver) Read(p []byte) (n int, err error) {
|
||||
packet, ok := <-r.readCh
|
||||
if ok {
|
||||
n = copy(p, packet)
|
||||
pool4K.Put(packet[:cap(packet)])
|
||||
return n, nil
|
||||
}
|
||||
return 0, io.ErrClosedPipe
|
||||
}
|
||||
|
||||
func (r *TCPResolver) Send(p []byte) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if r.dial() != nil {
|
||||
return
|
||||
}
|
||||
_ = binary.Write(r.conn, binary.BigEndian, len(p))
|
||||
_, _ = r.conn.Write(p)
|
||||
}
|
||||
|
||||
func (r *TCPResolver) Close() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if r.closed() {
|
||||
return
|
||||
}
|
||||
close(r.closeCh)
|
||||
if r.conn != nil {
|
||||
_ = r.conn.Close()
|
||||
}
|
||||
r.wg.Wait()
|
||||
close(r.readCh)
|
||||
}
|
||||
@@ -0,0 +1,130 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"sync"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
type UDPResolver struct {
|
||||
dest net.Destination
|
||||
dialer *finalmask.Dialer
|
||||
|
||||
conn net.PacketConn
|
||||
udpAddr *net.UDPAddr
|
||||
|
||||
readCh chan []byte
|
||||
closeCh chan struct{}
|
||||
wg sync.WaitGroup
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewUDPResolver(config *UDPResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
dest, err := net.ParseDestination("udp:" + config.Addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r := &UDPResolver{
|
||||
dest: dest,
|
||||
dialer: dialer,
|
||||
readCh: make(chan []byte),
|
||||
closeCh: make(chan struct{}),
|
||||
}
|
||||
if err := r.dial(); err != nil {
|
||||
r.Close()
|
||||
return nil, err
|
||||
}
|
||||
return r, nil
|
||||
}
|
||||
|
||||
func (r *UDPResolver) closed() bool {
|
||||
select {
|
||||
case <-r.closeCh:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (r *UDPResolver) dial() error {
|
||||
if r.closed() {
|
||||
return errors.New("closed")
|
||||
}
|
||||
if r.conn != nil {
|
||||
return nil
|
||||
}
|
||||
conn, err := r.dialer.DialUDP(r.dest)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.conn = conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
r.udpAddr = conn.RemoteAddr().(*net.UDPAddr)
|
||||
r.wg.Add(1)
|
||||
go r.recv(conn.(*finalmask.PacketConnWrapper).PacketConn)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *UDPResolver) recv(conn net.PacketConn) {
|
||||
defer r.wg.Done()
|
||||
|
||||
var buf [4096]byte
|
||||
for {
|
||||
n, _, err := conn.ReadFrom(buf[:])
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
p := pool4K.Get().([]byte)
|
||||
copy(p, buf[:n])
|
||||
select {
|
||||
case <-r.closeCh:
|
||||
pool4K.Put(p[:cap(p)])
|
||||
case r.readCh <- p[:n]:
|
||||
}
|
||||
}
|
||||
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
_ = conn.Close()
|
||||
r.conn = nil
|
||||
}
|
||||
|
||||
func (r *UDPResolver) Addr() *net.UDPAddr {
|
||||
return r.udpAddr
|
||||
}
|
||||
|
||||
func (r *UDPResolver) Read(p []byte) (n int, err error) {
|
||||
packet, ok := <-r.readCh
|
||||
if ok {
|
||||
n = copy(p, packet)
|
||||
pool4K.Put(packet[:cap(packet)])
|
||||
return n, nil
|
||||
}
|
||||
return 0, io.ErrClosedPipe
|
||||
}
|
||||
|
||||
func (r *UDPResolver) Send(p []byte) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if err := r.dial(); err != nil {
|
||||
return
|
||||
}
|
||||
_, _ = r.conn.WriteTo(p, r.udpAddr)
|
||||
}
|
||||
|
||||
func (r *UDPResolver) Close() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if r.closed() {
|
||||
return
|
||||
}
|
||||
close(r.closeCh)
|
||||
if r.conn != nil {
|
||||
_ = r.conn.Close()
|
||||
}
|
||||
r.wg.Wait()
|
||||
close(r.readCh)
|
||||
}
|
||||
@@ -0,0 +1,392 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
"golang.org/x/net/dns/dnsmessage"
|
||||
)
|
||||
|
||||
const (
|
||||
sendTTL = 4 * time.Second
|
||||
)
|
||||
|
||||
type Resp struct {
|
||||
msg dnsmessage.Message
|
||||
domain *Domain
|
||||
edns0 uint16
|
||||
|
||||
cap int
|
||||
}
|
||||
|
||||
func NewResp(msg dnsmessage.Message, domain *Domain, edns0 uint16) *Resp {
|
||||
if msg.Header.Response {
|
||||
return &Resp{
|
||||
msg: msg,
|
||||
domain: domain,
|
||||
}
|
||||
}
|
||||
|
||||
size := min(max(int(edns0), 512), max(int(domain.edns0), 512))
|
||||
|
||||
left := size - 12 - int(msg.Questions[0].Name.Length) - 1 - 2 - 2
|
||||
if edns0 > 0 {
|
||||
left -= 1 + 2 + 2 + 4 + 2 + 0
|
||||
}
|
||||
cap := 0
|
||||
switch msg.Questions[0].Type {
|
||||
case dnsmessage.TypeA:
|
||||
single := 2 + 2 + 2 + 4 + 2 + 4
|
||||
n := left / single
|
||||
if n > 255 {
|
||||
n = 255
|
||||
}
|
||||
cap = 4*n - n - 1
|
||||
case dnsmessage.TypeCNAME:
|
||||
single := 2 + 2 + 2 + 4 + 2 + domain.lenMax
|
||||
n := left / single
|
||||
if n > 255 {
|
||||
n = 255
|
||||
}
|
||||
cap = domain.cap*n - n - 1
|
||||
case dnsmessage.TypeTXT:
|
||||
left -= 2 + 2 + 2 + 4 + 2
|
||||
single := 255
|
||||
n := left / single
|
||||
m := left % single
|
||||
cap = 255*n - n
|
||||
if m > 1 {
|
||||
cap += m - 1
|
||||
}
|
||||
case dnsmessage.TypeAAAA:
|
||||
single := 2 + 2 + 2 + 4 + 2 + 16
|
||||
n := left / single
|
||||
if n > 255 {
|
||||
n = 255
|
||||
}
|
||||
cap = 16*n - n - 1
|
||||
}
|
||||
|
||||
return &Resp{
|
||||
msg: msg,
|
||||
domain: domain,
|
||||
edns0: edns0,
|
||||
|
||||
cap: cap,
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Resp) Encode(encoded []byte, data []byte) []byte {
|
||||
msg := r.msg
|
||||
msg.Header = dnsmessage.Header{
|
||||
ID: msg.Header.ID,
|
||||
Response: true,
|
||||
Authoritative: true,
|
||||
RCode: dnsmessage.RCodeSuccess,
|
||||
}
|
||||
msg.Answers = nil
|
||||
msg.Authorities = nil
|
||||
msg.Additionals = nil
|
||||
switch msg.Questions[0].Type {
|
||||
case dnsmessage.TypeA:
|
||||
fragN := 0
|
||||
if len(data) > 0 {
|
||||
fragN = 1
|
||||
}
|
||||
if (len(data) - (4 - 2)) > 0 {
|
||||
fragN += (len(data) - (4 - 2)) / (4 - 1)
|
||||
if (len(data)-(4-2))%(4-1) > 0 {
|
||||
fragN++
|
||||
}
|
||||
}
|
||||
|
||||
for i := range fragN {
|
||||
A := [4]byte{byte(i)}
|
||||
if i == 0 {
|
||||
A[1] = byte(fragN)
|
||||
n := copy(A[2:], data)
|
||||
data = data[n:]
|
||||
} else {
|
||||
n := copy(A[1:], data)
|
||||
data = data[n:]
|
||||
}
|
||||
msg.Answers = append(msg.Answers, dnsmessage.Resource{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: msg.Questions[0].Name,
|
||||
Type: msg.Questions[0].Type,
|
||||
Class: dnsmessage.ClassINET,
|
||||
TTL: 60,
|
||||
},
|
||||
Body: &dnsmessage.AResource{A: A},
|
||||
})
|
||||
}
|
||||
case dnsmessage.TypeCNAME:
|
||||
fragN := 0
|
||||
if len(data) > 0 {
|
||||
fragN = 1
|
||||
}
|
||||
if (len(data) - (r.domain.cap - 2)) > 0 {
|
||||
fragN += (len(data) - (r.domain.cap - 2)) / (r.domain.cap - 1)
|
||||
if (len(data)-(r.domain.cap-2))%(r.domain.cap-1) > 0 {
|
||||
fragN++
|
||||
}
|
||||
}
|
||||
|
||||
DATA := make([]byte, r.domain.cap)
|
||||
for i := range fragN {
|
||||
DATA[0] = byte(i)
|
||||
if i == 0 {
|
||||
DATA[1] = byte(fragN)
|
||||
n := copy(DATA[2:], data)
|
||||
data = data[n:]
|
||||
msg.Answers = append(msg.Answers, dnsmessage.Resource{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: msg.Questions[0].Name,
|
||||
Type: msg.Questions[0].Type,
|
||||
Class: dnsmessage.ClassINET,
|
||||
TTL: 60,
|
||||
},
|
||||
Body: &dnsmessage.CNAMEResource{CNAME: r.domain.Encode(DATA[:2+n])},
|
||||
})
|
||||
} else {
|
||||
n := copy(DATA[1:], data)
|
||||
data = data[n:]
|
||||
msg.Answers = append(msg.Answers, dnsmessage.Resource{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: msg.Questions[0].Name,
|
||||
Type: msg.Questions[0].Type,
|
||||
Class: dnsmessage.ClassINET,
|
||||
TTL: 60,
|
||||
},
|
||||
Body: &dnsmessage.CNAMEResource{CNAME: r.domain.Encode(DATA[:1+n])},
|
||||
})
|
||||
}
|
||||
}
|
||||
case dnsmessage.TypeTXT:
|
||||
var txt []string
|
||||
for len(data) > 0 {
|
||||
size := min(len(data), 255)
|
||||
txt = append(txt, string(data[:size]))
|
||||
data = data[size:]
|
||||
}
|
||||
msg.Answers = append(msg.Answers, dnsmessage.Resource{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: msg.Questions[0].Name,
|
||||
Type: msg.Questions[0].Type,
|
||||
Class: dnsmessage.ClassINET,
|
||||
TTL: 60,
|
||||
},
|
||||
Body: &dnsmessage.TXTResource{TXT: txt},
|
||||
})
|
||||
case dnsmessage.TypeAAAA:
|
||||
fragN := 0
|
||||
if len(data) > 0 {
|
||||
fragN = 1
|
||||
}
|
||||
if (len(data) - (16 - 2)) > 0 {
|
||||
fragN += (len(data) - (16 - 2)) / (16 - 1)
|
||||
if (len(data)-(16-2))%(16-1) > 0 {
|
||||
fragN++
|
||||
}
|
||||
}
|
||||
|
||||
for i := range fragN {
|
||||
AAAA := [16]byte{byte(i)}
|
||||
if i == 0 {
|
||||
AAAA[1] = byte(fragN)
|
||||
n := copy(AAAA[2:], data)
|
||||
data = data[n:]
|
||||
} else {
|
||||
n := copy(AAAA[1:], data)
|
||||
data = data[n:]
|
||||
}
|
||||
msg.Answers = append(msg.Answers, dnsmessage.Resource{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: msg.Questions[0].Name,
|
||||
Type: msg.Questions[0].Type,
|
||||
Class: dnsmessage.ClassINET,
|
||||
TTL: 60,
|
||||
},
|
||||
Body: &dnsmessage.AAAAResource{AAAA: AAAA},
|
||||
})
|
||||
}
|
||||
}
|
||||
if r.edns0 > 0 {
|
||||
msg.Additionals = append(msg.Additionals, dnsmessage.Resource{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: dnsmessage.MustNewName("."),
|
||||
Type: dnsmessage.TypeOPT,
|
||||
Class: dnsmessage.Class(r.edns0),
|
||||
TTL: 0,
|
||||
},
|
||||
Body: &dnsmessage.OPTResource{},
|
||||
})
|
||||
}
|
||||
return common.Must2(msg.AppendPack(encoded[:0]))
|
||||
}
|
||||
|
||||
func (r *Resp) Decode(decoded []byte) int {
|
||||
decoded = decoded[:0]
|
||||
msg := r.msg
|
||||
if msg.Questions[0].Type == dnsmessage.TypeTXT {
|
||||
if len(msg.Answers) == 1 && r.domain.IsDomain(msg.Answers[0].Header.Name) && msg.Answers[0].Header.Type == dnsmessage.TypeTXT {
|
||||
for i := range msg.Answers[0].Body.(*dnsmessage.TXTResource).TXT {
|
||||
decoded = append(decoded, msg.Answers[0].Body.(*dnsmessage.TXTResource).TXT[i]...)
|
||||
}
|
||||
}
|
||||
return len(decoded)
|
||||
} else {
|
||||
var frags [][]byte
|
||||
for i := range msg.Answers {
|
||||
if !r.domain.IsDomain(msg.Answers[i].Header.Name) || msg.Answers[i].Header.Type != msg.Questions[0].Type {
|
||||
continue
|
||||
}
|
||||
switch msg.Questions[0].Type {
|
||||
case dnsmessage.TypeA:
|
||||
frags = append(frags, msg.Answers[i].Body.(*dnsmessage.AResource).A[:])
|
||||
case dnsmessage.TypeCNAME:
|
||||
var decoded [255]byte
|
||||
n := r.domain.Decode(&decoded, msg.Answers[i].Body.(*dnsmessage.CNAMEResource).CNAME)
|
||||
if n == 0 {
|
||||
continue
|
||||
}
|
||||
frags = append(frags, decoded[:n])
|
||||
case dnsmessage.TypeAAAA:
|
||||
frags = append(frags, msg.Answers[i].Body.(*dnsmessage.AAAAResource).AAAA[:])
|
||||
}
|
||||
}
|
||||
sort.Slice(frags, func(i, j int) bool {
|
||||
return frags[i][0] < frags[j][0]
|
||||
})
|
||||
if len(frags) < 1 || len(frags[0]) < 2 || int(frags[0][1]) > len(frags) {
|
||||
return 0
|
||||
}
|
||||
decoded = append(decoded, frags[0][2:]...)
|
||||
for i := range frags {
|
||||
if i > 0 {
|
||||
if frags[i][0] == frags[i-1][0] {
|
||||
return 0
|
||||
}
|
||||
decoded = append(decoded, frags[i][1:]...)
|
||||
}
|
||||
}
|
||||
return len(decoded)
|
||||
}
|
||||
}
|
||||
|
||||
type SendInfo struct {
|
||||
stash chan []byte
|
||||
ch chan []byte
|
||||
deadline time.Time
|
||||
}
|
||||
|
||||
type SendManager struct {
|
||||
m map[ClientID]*SendInfo
|
||||
ch chan struct{}
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewSendManager() *SendManager {
|
||||
m := &SendManager{
|
||||
m: make(map[ClientID]*SendInfo),
|
||||
ch: make(chan struct{}),
|
||||
}
|
||||
go m.gc()
|
||||
return m
|
||||
}
|
||||
|
||||
func (m *SendManager) closed() bool {
|
||||
select {
|
||||
case <-m.ch:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (m *SendManager) gc() {
|
||||
ticker := time.NewTicker(sendTTL)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-m.ch:
|
||||
return
|
||||
case now := <-ticker.C:
|
||||
m.mu.Lock()
|
||||
for key, info := range m.m {
|
||||
if now.After(info.deadline) {
|
||||
close(info.stash)
|
||||
close(info.ch)
|
||||
delete(m.m, key)
|
||||
}
|
||||
}
|
||||
m.mu.Unlock()
|
||||
ticker.Reset(sendTTL)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m *SendManager) Push(clientID ClientID, p []byte) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
info := m.m[clientID]
|
||||
if info == nil {
|
||||
info = &SendInfo{
|
||||
stash: make(chan []byte, 1),
|
||||
ch: make(chan []byte, 128),
|
||||
deadline: time.Now().Add(sendTTL),
|
||||
}
|
||||
m.m[clientID] = info
|
||||
}
|
||||
b := make([]byte, len(p))
|
||||
copy(b, p)
|
||||
select {
|
||||
case info.ch <- b:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (m *SendManager) Stash(clientID ClientID, p []byte) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
info := m.m[clientID]
|
||||
if info == nil {
|
||||
return
|
||||
}
|
||||
info.deadline = time.Now().Add(sendTTL)
|
||||
select {
|
||||
case info.stash <- p:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (m *SendManager) Pop(clientID ClientID) (chan []byte, chan []byte) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
info := m.m[clientID]
|
||||
if info == nil {
|
||||
info = &SendInfo{
|
||||
stash: make(chan []byte, 1),
|
||||
ch: make(chan []byte, 128),
|
||||
}
|
||||
m.m[clientID] = info
|
||||
}
|
||||
info.deadline = time.Now().Add(sendTTL)
|
||||
return info.ch, info.stash
|
||||
}
|
||||
|
||||
func (m *SendManager) Close() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.closed() {
|
||||
return
|
||||
}
|
||||
close(m.ch)
|
||||
for key, info := range m.m {
|
||||
close(info.stash)
|
||||
close(info.ch)
|
||||
delete(m.m, key)
|
||||
}
|
||||
}
|
||||
@@ -1,512 +1,385 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
go_errors "errors"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"golang.org/x/net/dns/dnsmessage"
|
||||
)
|
||||
|
||||
const (
|
||||
idleTimeout = 10 * time.Second
|
||||
responseTTL = 60
|
||||
maxResponseDelay = 1 * time.Second
|
||||
maxResponseDelay = time.Second
|
||||
)
|
||||
|
||||
var (
|
||||
maxUDPPayload = 1280 - 40 - 8
|
||||
maxEncodedPayloadTXT = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeTXT)
|
||||
maxEncodedPayloadA = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeA)
|
||||
maxEncodedPayloadAAAA = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeAAAA)
|
||||
)
|
||||
|
||||
func clientIDToAddr(clientID [8]byte) *net.UDPAddr {
|
||||
ip := make(net.IP, 16)
|
||||
|
||||
copy(ip, []byte{0xfd, 0x00, 0, 0, 0, 0, 0, 0})
|
||||
copy(ip[8:], clientID[:])
|
||||
|
||||
return &net.UDPAddr{
|
||||
IP: ip,
|
||||
}
|
||||
type resp struct {
|
||||
msg dnsmessage.Message
|
||||
addr net.Addr
|
||||
}
|
||||
|
||||
type record struct {
|
||||
Resp *Message
|
||||
Addr net.Addr
|
||||
// ClientID [8]byte
|
||||
ClientAddr net.Addr
|
||||
type Rec struct {
|
||||
resp *Resp
|
||||
clientID ClientID
|
||||
addr net.Addr
|
||||
}
|
||||
|
||||
type queue struct {
|
||||
last time.Time
|
||||
rrType uint16
|
||||
queue chan []byte
|
||||
stash chan []byte
|
||||
}
|
||||
|
||||
type xdnsConnServer struct {
|
||||
type xdnsServer struct {
|
||||
net.PacketConn
|
||||
|
||||
domains []domainSpec
|
||||
domains []*Domain
|
||||
fragManager *FragManager
|
||||
sendManager *SendManager
|
||||
|
||||
ch chan *record
|
||||
readQueue chan *packet
|
||||
writeQueueMap map[string]*queue
|
||||
|
||||
closed bool
|
||||
mutex sync.Mutex
|
||||
readCh chan packet
|
||||
recCh chan *Rec
|
||||
drCh chan resp
|
||||
closeCh chan struct{}
|
||||
wg sync.WaitGroup
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
func NewServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
if len(c.Domains) == 0 {
|
||||
return nil, errors.New("empty domains")
|
||||
}
|
||||
domains := make([]domainSpec, 0, len(c.Domains))
|
||||
for _, domain := range c.Domains {
|
||||
domain, err := parseDomainSpec(domain, "")
|
||||
domains := make([]*Domain, 0, len(c.Domains))
|
||||
for i := range c.Domains {
|
||||
types := make([]uint16, 0, len(c.Domains[i].Types))
|
||||
for j := range c.Domains[i].Types {
|
||||
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||
}
|
||||
domain, err := NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domains = append(domains, domain)
|
||||
}
|
||||
|
||||
conn := &xdnsConnServer{
|
||||
server := &xdnsServer{
|
||||
PacketConn: raw,
|
||||
|
||||
domains: domains,
|
||||
domains: domains,
|
||||
fragManager: NewFragManager(),
|
||||
sendManager: NewSendManager(),
|
||||
|
||||
ch: make(chan *record, 500),
|
||||
readQueue: make(chan *packet, 512),
|
||||
writeQueueMap: make(map[string]*queue),
|
||||
readCh: make(chan packet),
|
||||
recCh: make(chan *Rec, 255),
|
||||
drCh: make(chan resp),
|
||||
closeCh: make(chan struct{}),
|
||||
}
|
||||
|
||||
go conn.clean()
|
||||
go conn.recvLoop()
|
||||
go conn.sendLoop()
|
||||
|
||||
return conn, nil
|
||||
go server.run()
|
||||
return server, nil
|
||||
}
|
||||
|
||||
func (c *xdnsConnServer) clean() {
|
||||
f := func() bool {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
if c.closed {
|
||||
return true
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
|
||||
for key, q := range c.writeQueueMap {
|
||||
if now.Sub(q.last) >= idleTimeout {
|
||||
close(q.queue)
|
||||
close(q.stash)
|
||||
delete(c.writeQueueMap, key)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsServer) closed() bool {
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
|
||||
for {
|
||||
time.Sleep(idleTimeout / 2)
|
||||
if f() {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsConnServer) ensureQueue(addr net.Addr) *queue {
|
||||
if c.closed {
|
||||
return nil
|
||||
}
|
||||
|
||||
q, ok := c.writeQueueMap[addr.String()]
|
||||
if !ok {
|
||||
q = &queue{
|
||||
queue: make(chan []byte, 512),
|
||||
stash: make(chan []byte, 1),
|
||||
}
|
||||
c.writeQueueMap[addr.String()] = q
|
||||
}
|
||||
q.last = time.Now()
|
||||
|
||||
return q
|
||||
}
|
||||
|
||||
func (c *xdnsConnServer) stash(queue *queue, p []byte) {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
if c.closed {
|
||||
return
|
||||
}
|
||||
|
||||
func (c *xdnsServer) decref(msg dnsmessage.Message, addr net.Addr) {
|
||||
select {
|
||||
case queue.stash <- p:
|
||||
case c.drCh <- resp{msg: msg, addr: addr}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsConnServer) recvLoop() {
|
||||
var buf [finalmask.UDPSize]byte
|
||||
func (c *xdnsServer) read(buf []byte, addr net.Addr) {
|
||||
msg := dnsmessage.Message{}
|
||||
if err := msg.Unpack(buf); err != nil {
|
||||
return
|
||||
}
|
||||
if msg.Header.Response {
|
||||
return
|
||||
}
|
||||
|
||||
for {
|
||||
if c.closed {
|
||||
break
|
||||
}
|
||||
if msg.Header.OpCode != 0 {
|
||||
msg.Header.Response = true
|
||||
msg.Header.RCode = dnsmessage.RCodeNotImplemented
|
||||
c.decref(msg, addr)
|
||||
return
|
||||
}
|
||||
|
||||
n, addr, err := c.PacketConn.ReadFrom(buf[:])
|
||||
if err != nil {
|
||||
if go_errors.Is(err, net.ErrClosed) {
|
||||
break
|
||||
if len(msg.Questions) != 1 {
|
||||
msg.Header.Response = true
|
||||
msg.Header.RCode = dnsmessage.RCodeFormatError
|
||||
c.decref(msg, addr)
|
||||
return
|
||||
}
|
||||
|
||||
opt := false
|
||||
edns0 := uint16(0)
|
||||
for i := range msg.Additionals {
|
||||
if msg.Additionals[i].Header.Type == dnsmessage.TypeOPT {
|
||||
if opt {
|
||||
msg.Header.RCode = dnsmessage.RCodeFormatError
|
||||
c.decref(msg, addr)
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
query, err := MessageFromWireFormat(buf[:n])
|
||||
if err != nil {
|
||||
errors.LogDebug(context.Background(), addr, " xdns from wireformat err ", err)
|
||||
continue
|
||||
}
|
||||
|
||||
resp, payload := responseFor(&query, c.domains)
|
||||
|
||||
var clientID [8]byte
|
||||
n = copy(clientID[:], payload)
|
||||
payload = payload[n:]
|
||||
if n == len(clientID) {
|
||||
r := bytes.NewReader(payload)
|
||||
for {
|
||||
p, err := nextPacketServer(r)
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
|
||||
buf := make([]byte, len(p))
|
||||
copy(buf, p)
|
||||
select {
|
||||
case c.readQueue <- &packet{
|
||||
p: buf,
|
||||
addr: clientIDToAddr(clientID),
|
||||
}:
|
||||
default:
|
||||
errors.LogDebug(context.Background(), addr, " ", clientID, " mask read err queue full")
|
||||
}
|
||||
}
|
||||
} else {
|
||||
if resp != nil && resp.Rcode() == RcodeNoError {
|
||||
resp.Flags |= RcodeNameError
|
||||
}
|
||||
}
|
||||
|
||||
if resp != nil {
|
||||
select {
|
||||
case c.ch <- &record{resp, addr, clientIDToAddr(clientID)}:
|
||||
default:
|
||||
errors.LogDebug(context.Background(), addr, " ", clientID, " mask read err record queue full")
|
||||
opt = true
|
||||
edns0 = uint16(msg.Additionals[i].Header.Class)
|
||||
if ver := (msg.Additionals[i].Header.TTL >> 16) & 0xFF; ver != 0 {
|
||||
msg.Header.RCode = dnsmessage.RCodeSuccess
|
||||
msg.Additionals[i].Header.TTL = 1 << 24
|
||||
c.decref(msg, addr)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
if opt {
|
||||
if edns0 < 512 {
|
||||
edns0 = 512
|
||||
}
|
||||
if edns0 > 4096 {
|
||||
edns0 = 4096
|
||||
}
|
||||
}
|
||||
errors.LogDebug(context.Background(), addr, " edns0 ", edns0, " buf ", len(buf), " ", msg.Questions[0].Type)
|
||||
|
||||
errors.LogDebug(context.Background(), "xdns closed")
|
||||
var domain *Domain
|
||||
for i := range c.domains {
|
||||
if c.domains[i].IsDomain(msg.Questions[0].Name) {
|
||||
domain = c.domains[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
if domain == nil {
|
||||
msg.Header.Response = true
|
||||
msg.Header.RCode = dnsmessage.RCodeNameError
|
||||
c.decref(msg, addr)
|
||||
return
|
||||
}
|
||||
if !domain.HasType(uint16(msg.Questions[0].Type)) {
|
||||
msg.Header.Response = true
|
||||
msg.Header.Authoritative = true
|
||||
msg.Header.RCode = dnsmessage.RCodeSuccess
|
||||
c.decref(msg, addr)
|
||||
return
|
||||
}
|
||||
|
||||
close(c.ch)
|
||||
close(c.readQueue)
|
||||
var decoded [255]byte
|
||||
n := domain.Decode(&decoded, msg.Questions[0].Name)
|
||||
if n < 9 {
|
||||
msg.Header.Response = true
|
||||
msg.Header.Authoritative = true
|
||||
msg.Header.RCode = dnsmessage.RCodeSuccess
|
||||
c.decref(msg, addr)
|
||||
return
|
||||
}
|
||||
if TypeMap_[decoded[0]&3] != uint16(msg.Questions[0].Type) || (decoded[8]&0x3F != 3 && decoded[8]&0x3F != 8) || (decoded[8]&0x3F == 3 && n < 9+3+1) || (decoded[8]&0x3F == 8 && n != 9+8) {
|
||||
msg.Header.Response = true
|
||||
msg.Header.Authoritative = true
|
||||
msg.Header.RCode = dnsmessage.RCodeSuccess
|
||||
c.decref(msg, addr)
|
||||
return
|
||||
}
|
||||
clientID := ClientIDFromRaw([8]byte(decoded[:8]))
|
||||
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
r := NewResp(msg, domain, edns0)
|
||||
if r == nil {
|
||||
msg.Header.Response = true
|
||||
msg.Header.Authoritative = true
|
||||
msg.Header.RCode = dnsmessage.RCodeSuccess
|
||||
c.decref(msg, addr)
|
||||
return
|
||||
}
|
||||
select {
|
||||
case c.recCh <- &Rec{resp: r, clientID: clientID, addr: addr}:
|
||||
default:
|
||||
msg.Header.Response = true
|
||||
msg.Header.Authoritative = true
|
||||
msg.Header.RCode = dnsmessage.RCodeSuccess
|
||||
c.decref(msg, addr)
|
||||
}
|
||||
|
||||
c.closed = true
|
||||
for key, q := range c.writeQueueMap {
|
||||
close(q.queue)
|
||||
close(q.stash)
|
||||
delete(c.writeQueueMap, key)
|
||||
if decoded[8]&0x3F == 8 {
|
||||
return
|
||||
}
|
||||
p := pool4K.Get().([]byte)
|
||||
p = p[:0]
|
||||
if decoded[8]&0xC0 == 0xC0 {
|
||||
out := pool4K.Get().([]byte)
|
||||
n := c.fragManager.Feed(out, FragKey{clientID: clientID, fragID: decoded[12]}, decoded[13], decoded[14], decoded[15:n])
|
||||
pool4K.Put(p[:cap(p)])
|
||||
if n > 0 {
|
||||
p = out[:n]
|
||||
} else {
|
||||
pool4K.Put(out[:cap(out)])
|
||||
return
|
||||
}
|
||||
} else {
|
||||
p = append(p, decoded[12:n]...)
|
||||
}
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
pool4K.Put(p[:cap(p)])
|
||||
return
|
||||
case c.readCh <- packet{p: p, addr: clientID.Addr()}:
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsConnServer) sendLoop() {
|
||||
var nextRec *record
|
||||
func (c *xdnsServer) run() {
|
||||
c.wg.Add(1)
|
||||
go c.recv()
|
||||
|
||||
c.wg.Add(1)
|
||||
go c.send()
|
||||
|
||||
c.wg.Add(1)
|
||||
go c.dr()
|
||||
|
||||
c.wg.Wait()
|
||||
close(c.readCh)
|
||||
close(c.recCh)
|
||||
close(c.drCh)
|
||||
c.fragManager.Close()
|
||||
c.sendManager.Close()
|
||||
}
|
||||
|
||||
func (c *xdnsServer) recv() {
|
||||
defer c.wg.Done()
|
||||
|
||||
var buf [512]byte
|
||||
for {
|
||||
n, addr, err := c.PacketConn.ReadFrom(buf[:])
|
||||
if err != nil {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
errors.LogErrorInner(context.Background(), err, "recv err")
|
||||
return
|
||||
}
|
||||
c.read(buf[:n], addr)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsServer) send() {
|
||||
defer c.wg.Done()
|
||||
|
||||
timer := time.NewTimer(maxResponseDelay)
|
||||
timer.Stop()
|
||||
var buf [4096]byte
|
||||
var data [4096]byte
|
||||
var nextRec *Rec
|
||||
for {
|
||||
var err error
|
||||
rec := nextRec
|
||||
nextRec = nil
|
||||
|
||||
if rec == nil {
|
||||
var ok bool
|
||||
rec, ok = <-c.ch
|
||||
if !ok {
|
||||
break
|
||||
select {
|
||||
case rec = <-c.recCh:
|
||||
case <-c.closeCh:
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if rec.Resp.Rcode() == RcodeNoError && len(rec.Resp.Question) == 1 {
|
||||
var payload bytes.Buffer
|
||||
limit := maxEncodedPayloadForType(rec.Resp.Question[0].Type)
|
||||
timer := time.NewTimer(maxResponseDelay)
|
||||
|
||||
for {
|
||||
c.mutex.Lock()
|
||||
q := c.ensureQueue(rec.ClientAddr)
|
||||
if q == nil {
|
||||
c.mutex.Unlock()
|
||||
return
|
||||
}
|
||||
q.rrType = rec.Resp.Question[0].Type
|
||||
c.mutex.Unlock()
|
||||
|
||||
var p []byte
|
||||
|
||||
ch, stash := c.sendManager.Pop(rec.clientID)
|
||||
left := rec.resp.cap
|
||||
timer.Reset(maxResponseDelay)
|
||||
var ps [][]byte
|
||||
for {
|
||||
var p []byte
|
||||
select {
|
||||
case p = <-stash:
|
||||
default:
|
||||
select {
|
||||
case p = <-q.stash:
|
||||
case p = <-stash:
|
||||
case p = <-ch:
|
||||
default:
|
||||
select {
|
||||
case p = <-q.stash:
|
||||
case p = <-q.queue:
|
||||
default:
|
||||
select {
|
||||
case p = <-q.stash:
|
||||
case p = <-q.queue:
|
||||
case <-timer.C:
|
||||
case nextRec = <-c.ch:
|
||||
}
|
||||
case p = <-stash:
|
||||
case p = <-ch:
|
||||
case <-timer.C:
|
||||
case nextRec = <-c.recCh:
|
||||
}
|
||||
}
|
||||
|
||||
timer.Reset(0)
|
||||
|
||||
if len(p) == 0 {
|
||||
}
|
||||
if len(p) == 0 {
|
||||
break
|
||||
}
|
||||
timer.Reset(0)
|
||||
left -= 2 + len(p)
|
||||
if left < 0 {
|
||||
if len(ps) == 0 {
|
||||
errors.LogError(context.Background(), "err size ", len(p))
|
||||
break
|
||||
}
|
||||
|
||||
limit -= 2 + len(p)
|
||||
if limit < 0 {
|
||||
if payload.Len() == 0 {
|
||||
errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns payload too large for rrtype ", rec.Resp.Question[0].Type, " ", len(p))
|
||||
continue
|
||||
}
|
||||
c.stash(q, p)
|
||||
break
|
||||
}
|
||||
|
||||
// if len(p) > 65535 {
|
||||
// panic(len(p))
|
||||
// }
|
||||
|
||||
_ = binary.Write(&payload, binary.BigEndian, uint16(len(p)))
|
||||
payload.Write(p)
|
||||
c.sendManager.Stash(rec.clientID, p)
|
||||
break
|
||||
}
|
||||
ps = append(ps, p)
|
||||
}
|
||||
timer.Stop()
|
||||
|
||||
timer.Stop()
|
||||
rec.Resp.Answer, err = answersForPayload(rec.Resp.Question[0], responseTTL, payload.Bytes())
|
||||
if err != nil {
|
||||
errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns encode err ", err)
|
||||
continue
|
||||
d := data[:0]
|
||||
for i := range ps {
|
||||
l := len(ps[i])
|
||||
if i == len(ps)-1 {
|
||||
l |= 0xC000
|
||||
}
|
||||
d = append(d, []byte{byte(l >> 8), byte(l)}...)
|
||||
d = append(d, ps[i]...)
|
||||
}
|
||||
_, _ = c.PacketConn.WriteTo(rec.resp.Encode(buf[:0], d), rec.addr)
|
||||
}
|
||||
}
|
||||
|
||||
buf, err := rec.Resp.WireFormat()
|
||||
if err != nil {
|
||||
errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns wireformat err ", err)
|
||||
continue
|
||||
}
|
||||
func (c *xdnsServer) dr() {
|
||||
defer c.wg.Done()
|
||||
|
||||
if len(buf) > maxUDPPayload {
|
||||
errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns truncate ", len(buf))
|
||||
buf = buf[:maxUDPPayload]
|
||||
buf[2] |= 0x02
|
||||
}
|
||||
|
||||
if c.closed {
|
||||
var buf [512]byte
|
||||
for {
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
}
|
||||
|
||||
_, err = c.PacketConn.WriteTo(buf, rec.Addr)
|
||||
if go_errors.Is(err, net.ErrClosed) {
|
||||
c.closed = true
|
||||
break
|
||||
case r := <-c.drCh:
|
||||
_, _ = c.PacketConn.WriteTo(common.Must2(r.msg.AppendPack(buf[:0])), r.addr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsConnServer) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
||||
packet, ok := <-c.readQueue
|
||||
if !ok {
|
||||
return 0, nil, net.ErrClosed
|
||||
func (c *xdnsServer) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
||||
packet, ok := <-c.readCh
|
||||
if ok {
|
||||
n = copy(p, packet.p)
|
||||
pool4K.Put(packet.p[:cap(packet.p)])
|
||||
return n, packet.addr, nil
|
||||
}
|
||||
if len(p) < len(packet.p) {
|
||||
errors.LogDebug(context.Background(), packet.addr, " mask read err short buffer ", len(p), " ", len(packet.p))
|
||||
return 0, packet.addr, nil
|
||||
}
|
||||
copy(p, packet.p)
|
||||
return len(packet.p), packet.addr, nil
|
||||
return 0, nil, io.ErrClosedPipe
|
||||
}
|
||||
|
||||
func (c *xdnsConnServer) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
q := c.ensureQueue(addr)
|
||||
if q == nil {
|
||||
func (c *xdnsServer) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
if c.closed() {
|
||||
return 0, io.ErrClosedPipe
|
||||
}
|
||||
limit := maxEncodedPayloadForType(q.rrType)
|
||||
if q.rrType == 0 {
|
||||
limit = maxEncodedPayloadTXT
|
||||
}
|
||||
if len(p)+2 > limit {
|
||||
errors.LogDebug(context.Background(), addr, " mask write err short write ", len(p), "+2 > ", limit)
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
buf := make([]byte, len(p))
|
||||
copy(buf, p)
|
||||
|
||||
select {
|
||||
case q.queue <- buf:
|
||||
return len(p), nil
|
||||
default:
|
||||
// errors.LogDebug(context.Background(), addr, " mask write err queue full")
|
||||
return 0, nil
|
||||
if len(p) == 0 || len(p) > 4096 {
|
||||
errors.LogError(context.Background(), "err size ", len(p))
|
||||
return 0, errors.New("err size")
|
||||
}
|
||||
c.sendManager.Push(ClientIDFromAddr(addr.(*net.UDPAddr)), p)
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func (c *xdnsConnServer) Close() error {
|
||||
c.closed = true
|
||||
return c.PacketConn.Close()
|
||||
func (c *xdnsServer) Close() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed() {
|
||||
return nil
|
||||
}
|
||||
close(c.closeCh)
|
||||
_ = c.PacketConn.Close()
|
||||
return nil
|
||||
}
|
||||
|
||||
func nextPacketServer(r *bytes.Reader) ([]byte, error) {
|
||||
eof := func(err error) error {
|
||||
if err == io.EOF {
|
||||
err = io.ErrUnexpectedEOF
|
||||
}
|
||||
return err
|
||||
}
|
||||
func (c *xdnsServer) SetDeadline(t time.Time) error { return errors.New("not support") }
|
||||
|
||||
for {
|
||||
prefix, err := r.ReadByte()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if prefix >= 224 {
|
||||
paddingLen := prefix - 224
|
||||
_, err := io.CopyN(io.Discard, r, int64(paddingLen))
|
||||
if err != nil {
|
||||
return nil, eof(err)
|
||||
}
|
||||
} else {
|
||||
p := make([]byte, int(prefix))
|
||||
_, err = io.ReadFull(r, p)
|
||||
return p, eof(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
func (c *xdnsServer) SetReadDeadline(t time.Time) error { return errors.New("not support") }
|
||||
|
||||
func responseFor(query *Message, domains []domainSpec) (*Message, []byte) {
|
||||
resp := &Message{
|
||||
ID: query.ID,
|
||||
Flags: 0x8000,
|
||||
Question: query.Question,
|
||||
}
|
||||
|
||||
if query.Flags&0x8000 != 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
payloadSize := 0
|
||||
for _, rr := range query.Additional {
|
||||
if rr.Type != RRTypeOPT {
|
||||
continue
|
||||
}
|
||||
if len(resp.Additional) != 0 {
|
||||
resp.Flags |= RcodeFormatError
|
||||
return resp, nil
|
||||
}
|
||||
resp.Additional = append(resp.Additional, RR{
|
||||
Name: Name{},
|
||||
Type: RRTypeOPT,
|
||||
Class: 4096,
|
||||
TTL: 0,
|
||||
Data: []byte{},
|
||||
})
|
||||
additional := &resp.Additional[0]
|
||||
|
||||
version := (rr.TTL >> 16) & 0xff
|
||||
if version != 0 {
|
||||
resp.Flags |= ExtendedRcodeBadVers & 0xf
|
||||
additional.TTL = (ExtendedRcodeBadVers >> 4) << 24
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
payloadSize = int(rr.Class)
|
||||
}
|
||||
if payloadSize < 512 {
|
||||
payloadSize = 512
|
||||
}
|
||||
|
||||
if len(query.Question) != 1 {
|
||||
resp.Flags |= RcodeFormatError
|
||||
return resp, nil
|
||||
}
|
||||
question := query.Question[0]
|
||||
|
||||
var (
|
||||
prefix Name
|
||||
ok bool
|
||||
match domainSpec
|
||||
)
|
||||
for _, domain := range domains {
|
||||
prefix, ok = question.Name.TrimSuffix(domain.name)
|
||||
if ok {
|
||||
match = domain
|
||||
break
|
||||
}
|
||||
}
|
||||
if !ok {
|
||||
resp.Flags |= RcodeNameError
|
||||
return resp, nil
|
||||
}
|
||||
resp.Flags |= 0x0400
|
||||
|
||||
if query.Opcode() != 0 {
|
||||
resp.Flags |= RcodeNotImplemented
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
switch question.Type {
|
||||
case RRTypeTXT, RRTypeA, RRTypeAAAA:
|
||||
default:
|
||||
resp.Flags |= RcodeNameError
|
||||
return resp, nil
|
||||
}
|
||||
if match.rrType != 0 && question.Type != match.rrType {
|
||||
resp.Flags |= RcodeNameError
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
encoded := bytes.ToUpper(bytes.Join(prefix, nil))
|
||||
payload := make([]byte, base32Encoding.DecodedLen(len(encoded)))
|
||||
n, err := base32Encoding.Decode(payload, encoded)
|
||||
if err != nil {
|
||||
resp.Flags |= RcodeNameError
|
||||
return resp, nil
|
||||
}
|
||||
payload = payload[:n]
|
||||
|
||||
if payloadSize < maxUDPPayload {
|
||||
resp.Flags |= RcodeFormatError
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
return resp, payload
|
||||
}
|
||||
func (c *xdnsServer) SetWriteDeadline(t time.Time) error { return errors.New("not support") }
|
||||
|
||||
@@ -1,80 +0,0 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
)
|
||||
|
||||
type domainSpec struct {
|
||||
name Name
|
||||
rrType uint16
|
||||
}
|
||||
|
||||
func rrTypeFromMethod(method string) (uint16, error) {
|
||||
switch strings.ToLower(method) {
|
||||
case "", "txt":
|
||||
return RRTypeTXT, nil
|
||||
case "a":
|
||||
return RRTypeA, nil
|
||||
case "aaaa":
|
||||
return RRTypeAAAA, nil
|
||||
default:
|
||||
return 0, errors.New("unsupported method")
|
||||
}
|
||||
}
|
||||
|
||||
func parseDomainSpec(s string, defaultMethod string) (domainSpec, error) {
|
||||
domainPart := s
|
||||
method := ""
|
||||
hasMethod := false
|
||||
|
||||
if i := strings.LastIndex(s, ":"); i >= 0 {
|
||||
domainPart = s[:i]
|
||||
method = s[i+1:]
|
||||
hasMethod = true
|
||||
} else if defaultMethod != "" {
|
||||
method = defaultMethod
|
||||
hasMethod = true
|
||||
}
|
||||
|
||||
if domainPart == "" {
|
||||
return domainSpec{}, errors.New("empty domain")
|
||||
}
|
||||
|
||||
name, err := ParseName(domainPart)
|
||||
if err != nil {
|
||||
return domainSpec{}, err
|
||||
}
|
||||
|
||||
rrType := uint16(0)
|
||||
if hasMethod {
|
||||
var err error
|
||||
rrType, err = rrTypeFromMethod(method)
|
||||
if err != nil {
|
||||
return domainSpec{}, err
|
||||
}
|
||||
}
|
||||
|
||||
return domainSpec{
|
||||
name: name,
|
||||
rrType: rrType,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func parseResolver(s string) (Name, string, uint16, error) {
|
||||
head, server, ok := strings.Cut(s, "+udp://")
|
||||
if !ok {
|
||||
return nil, "", 0, errors.New("invalid resolver scheme")
|
||||
}
|
||||
if server == "" {
|
||||
return nil, "", 0, errors.New("empty resolver server")
|
||||
}
|
||||
|
||||
spec, err := parseDomainSpec(head, "txt")
|
||||
if err != nil {
|
||||
return nil, "", 0, err
|
||||
}
|
||||
|
||||
return spec.name, server, spec.rrType, nil
|
||||
}
|
||||
@@ -0,0 +1,208 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/rand"
|
||||
"fmt"
|
||||
mrand "math/rand"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
"golang.org/x/net/dns/dnsmessage"
|
||||
)
|
||||
|
||||
func TestXxx(t *testing.T) {
|
||||
m1 := dnsmessage.Message{
|
||||
Questions: []dnsmessage.Question{
|
||||
{
|
||||
Name: dnsmessage.MustNewName("a.example.com."),
|
||||
},
|
||||
},
|
||||
Answers: []dnsmessage.Resource{
|
||||
{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: dnsmessage.MustNewName("a.example.com."),
|
||||
Type: dnsmessage.TypeA,
|
||||
Class: dnsmessage.ClassINET,
|
||||
TTL: 60,
|
||||
Length: 16,
|
||||
},
|
||||
Body: &dnsmessage.AResource{A: [4]byte{127, 0, 0, 1}},
|
||||
},
|
||||
},
|
||||
Additionals: []dnsmessage.Resource{
|
||||
{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: dnsmessage.MustNewName("."),
|
||||
Type: dnsmessage.TypeOPT,
|
||||
Class: 255,
|
||||
TTL: 0,
|
||||
Length: 16,
|
||||
},
|
||||
Body: &dnsmessage.OPTResource{},
|
||||
},
|
||||
},
|
||||
}
|
||||
p1, e1 := m1.Pack()
|
||||
if e1 != nil {
|
||||
t.Fatal(e1)
|
||||
}
|
||||
if !bytes.Equal(p1, []byte{
|
||||
0, 0, 0, 0, 0, 1, 0, 1, 0, 0, 0, 1,
|
||||
1, 97, 7, 101, 120, 97, 109, 112, 108, 101, 3, 99, 111, 109, 0,
|
||||
0, 0,
|
||||
0, 0,
|
||||
192, 12,
|
||||
0, 1,
|
||||
0, 1,
|
||||
0, 0, 0, 60,
|
||||
0, 4,
|
||||
127, 0, 0, 1,
|
||||
0,
|
||||
0, 41,
|
||||
0, 255,
|
||||
0, 0, 0, 0,
|
||||
0, 0,
|
||||
}) {
|
||||
t.Fatal("!bytes.Equal")
|
||||
}
|
||||
|
||||
domain, _ := NewDomain("a.example.com", 200, 1, []uint16{1}, 0)
|
||||
fmt.Println(domain.cap, domain.lenMax)
|
||||
lenMax := domain.lenMax
|
||||
data := make([]byte, domain.cap)
|
||||
msg := dnsmessage.Message{}
|
||||
msg.Unpack(p1)
|
||||
for range 3 {
|
||||
msg.Answers = nil
|
||||
msg.Authorities = nil
|
||||
msg.Additionals = nil
|
||||
n := mrand.Intn(255)
|
||||
for range n {
|
||||
msg.Answers = append(msg.Answers, dnsmessage.Resource{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: dnsmessage.MustNewName("a.example.com."),
|
||||
Type: dnsmessage.TypeA,
|
||||
Class: dnsmessage.ClassINET,
|
||||
TTL: 60,
|
||||
},
|
||||
Body: &dnsmessage.AResource{A: [4]byte{127, 0, 0, 1}},
|
||||
})
|
||||
}
|
||||
if len(common.Must2(msg.Pack())) != 12+15+2+2+n*(2+2+2+4+2+4) {
|
||||
t.Fatal("fatal a")
|
||||
}
|
||||
}
|
||||
for range 3 {
|
||||
msg.Answers = nil
|
||||
msg.Authorities = nil
|
||||
msg.Additionals = nil
|
||||
n := mrand.Intn(255)
|
||||
for range n {
|
||||
common.Must2(rand.Read(data))
|
||||
msg.Answers = append(msg.Answers, dnsmessage.Resource{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: dnsmessage.MustNewName("a.example.com."),
|
||||
Type: dnsmessage.TypeCNAME,
|
||||
Class: dnsmessage.ClassINET,
|
||||
TTL: 60,
|
||||
},
|
||||
Body: &dnsmessage.CNAMEResource{
|
||||
CNAME: domain.Encode(data),
|
||||
},
|
||||
})
|
||||
}
|
||||
if len(common.Must2(msg.Pack())) > 12+15+2+2+n*(2+2+2+4+2+lenMax) {
|
||||
t.Fatal("fatal cname")
|
||||
}
|
||||
}
|
||||
for range 3 {
|
||||
msg.Answers = nil
|
||||
msg.Authorities = nil
|
||||
msg.Additionals = nil
|
||||
n := (mrand.Intn(2048) + 1024) % 2048
|
||||
a := n / 255
|
||||
b := n % 255
|
||||
c := 0
|
||||
var d [255]byte
|
||||
var s []string
|
||||
for range a {
|
||||
s = append(s, string(d[:]))
|
||||
}
|
||||
if b > 0 {
|
||||
c = 1
|
||||
s = append(s, string(d[:b]))
|
||||
}
|
||||
msg.Answers = append(msg.Answers, dnsmessage.Resource{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: dnsmessage.MustNewName("a.example.com."),
|
||||
Type: dnsmessage.TypeTXT,
|
||||
Class: dnsmessage.ClassINET,
|
||||
TTL: 60,
|
||||
},
|
||||
Body: &dnsmessage.TXTResource{TXT: s},
|
||||
})
|
||||
if len(common.Must2(msg.Pack())) != 12+15+2+2+(2+2+2+4+2+n+n/255+c) {
|
||||
t.Fatal("fatal txt")
|
||||
}
|
||||
}
|
||||
for range 3 {
|
||||
msg.Answers = nil
|
||||
msg.Authorities = nil
|
||||
msg.Additionals = nil
|
||||
n := mrand.Intn(255)
|
||||
for range n {
|
||||
msg.Answers = append(msg.Answers, dnsmessage.Resource{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: dnsmessage.MustNewName("a.example.com."),
|
||||
Type: dnsmessage.TypeAAAA,
|
||||
Class: dnsmessage.ClassINET,
|
||||
TTL: 60,
|
||||
},
|
||||
Body: &dnsmessage.AAAAResource{AAAA: [16]byte{}},
|
||||
})
|
||||
}
|
||||
if len(common.Must2(msg.Pack())) != 12+15+2+2+n*(2+2+2+4+2+16) {
|
||||
t.Fatal("fatal aaaa")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTXT(t *testing.T) {
|
||||
txt := [][]byte{{}, {}}
|
||||
for i := range 255 {
|
||||
txt[0] = append(txt[0], byte(i))
|
||||
}
|
||||
txt[1] = []byte{255}
|
||||
str := []string{}
|
||||
for i := range txt {
|
||||
str = append(str, string(txt[i]))
|
||||
}
|
||||
m1 := dnsmessage.Message{
|
||||
Answers: []dnsmessage.Resource{
|
||||
{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: dnsmessage.MustNewName("."),
|
||||
Type: dnsmessage.TypeTXT,
|
||||
Class: dnsmessage.ClassINET,
|
||||
TTL: 60,
|
||||
},
|
||||
Body: &dnsmessage.TXTResource{
|
||||
TXT: str,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
p1 := common.Must2(m1.Pack())
|
||||
|
||||
m2 := dnsmessage.Message{}
|
||||
common.Must(m2.Unpack(p1))
|
||||
if len(m2.Answers[0].Body.(*dnsmessage.TXTResource).TXT) != len(txt) {
|
||||
t.Fatal("fatal txt")
|
||||
}
|
||||
for i := range txt {
|
||||
if !bytes.Equal(txt[i], []byte(m2.Answers[0].Body.(*dnsmessage.TXTResource).TXT[i])) {
|
||||
t.Fatal("fatal txt")
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -310,13 +310,6 @@ func (c *xicmpConnClient) Close() error {
|
||||
_ = c.icmp4.Close()
|
||||
_ = c.icmp6.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case p := <-c.readCh:
|
||||
if p.p != nil {
|
||||
pool.Put(p.p)
|
||||
}
|
||||
default:
|
||||
}
|
||||
close(c.readCh)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -329,13 +329,6 @@ func (c *xicmpConnServer) Close() error {
|
||||
_ = c.icmp4.Close()
|
||||
_ = c.icmp6.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case p := <-c.readCh:
|
||||
if p.p != nil {
|
||||
pool.Put(p.p)
|
||||
}
|
||||
default:
|
||||
}
|
||||
close(c.readCh)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -340,13 +340,6 @@ func (c *xicmpConnServer) Close() error {
|
||||
_ = c.icmp4.Close()
|
||||
_ = c.icmp6.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case p := <-c.readCh:
|
||||
if p.p != nil {
|
||||
pool.Put(p.p)
|
||||
}
|
||||
default:
|
||||
}
|
||||
close(c.readCh)
|
||||
return nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user