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:
LjhAUMEM
2026-09-29 17:16:59 +00:00
committed by GitHub
parent e5e85ca9da
commit fc8f8a451d
21 changed files with 2235 additions and 2647 deletions
+79 -21
View File
@@ -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
}
+344 -320
View File
@@ -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()
}
return nil
}
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)
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[:]}
}
+2 -2
View File
@@ -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)
}
+211 -19
View File
@@ -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,
},
+21 -2
View File
@@ -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;
}
-581
View File
@@ -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)
}
}
+215
View File
@@ -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
}
+171
View File
@@ -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)
}
+392
View File
@@ -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)
}
}
+288 -415
View File
@@ -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") }
-80
View File
@@ -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
}